新闻详情

新闻详情

首页 / 资讯中心 / 详情

实现24类花卉的高精度分类 PyTorch训练花卉分类数据集24类 使用花卉数据集进行图像分类

发布时间:2026/9/7 13:43:50来源:尧图网络
实现24类花卉的高精度分类 PyTorch训练花卉分类数据集24类 使用花卉数据集进行图像分类
采用预训练模型如ResNet进行实现24类花卉的高精度分类 PyTorch训练花卉分类数据集24类 使用花卉数据集进行图像分类以下文字及代码仅供参考学习使用。文章目录 1. 环境准备 2. 数据集结构要求 3. 数据加载器构建 4. 模型定义使用 ResNet50⚙️ 5. 训练配置️‍♂️ 6. 模型训练循环✅ 7. 测试评估 8. 可视化预测结果可选数据集描述**花卉数据集一共包含了47770张图片分为24类每一类包含了2500张图片图片的尺寸为224x224。具体分类为鬼针草、桔梗、石龙芮、全叶马兰、婆婆纳、三叶草、旋覆花、绣球小冠花、狗尾草、一年蓬、剑叶金鸡菊、滨菊、射干、三角梅、马鞭草、油菜花、蒲公英、两色金鸡菊、全缘金光菊、蓝蓟、曼陀罗、诸葛菜、千屈菜、狼尾草。适用于图像分类植物学分类中的花卉分类**使用花卉数据集进行图像分类的完整PyTorch训练代码。我们将采用预训练模型如ResNet进行微调以实现24类花卉的高精度分类。 1. 环境准备确保已安装以下依赖pipinstalltorch torchvision pandas matplotlib tqdm 2. 数据集结构要求你的数据集应按照如下格式组织flowers_dataset/ ├── train/ │ ├── class1/ │ ├── class2/ │ └── ... ├── val/ │ ├── class1/ │ ├── class2/ │ └── ... ├── test/ │ ├── class1/ │ ├── class2/ │ └── ... └── labels.txt其中labels.txt包含类别名称列表每行一个顺序与文件夹一致。每个子目录对应一个花卉种类包含2500张图片。 3. 数据加载器构建importosfromtorchvisionimporttransforms,datasetsfromtorch.utils.dataimportDataLoader# 数据增强和标准化transformtransforms.Compose([transforms.Resize((224,224)),transforms.ToTensor(),transforms.Normalize(mean[0.485,0.456,0.406],std[0.229,0.224,0.225])])# 数据集路径data_dirflowers_datasettrain_datasetdatasets.ImageFolder(os.path.join(data_dir,train),transformtransform)val_datasetdatasets.ImageFolder(os.path.join(data_dir,val),transformtransform)test_datasetdatasets.ImageFolder(os.path.join(data_dir,test),transformtransform)# DataLoaderbatch_size64train_loaderDataLoader(train_dataset,batch_sizebatch_size,shuffleTrue,num_workers4)val_loaderDataLoader(val_dataset,batch_sizebatch_size,shuffleFalse,num_workers4)test_loaderDataLoader(test_dataset,batch_sizebatch_size,shuffleFalse,num_workers4)print(Number of classes:,len(train_dataset.classes))print(Class names:,train_dataset.classes) 4. 模型定义使用 ResNet50importtorchimporttorch.nnasnnfromtorchvisionimportmodels devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)# 使用预训练的ResNet50modelmodels.resnet50(pretrainedTrue)# 修改最后一层全连接层适配24类num_ftrsmodel.fc.in_features model.fcnn.Linear(num_ftrs,24)# 24种花卉modelmodel.to(device)# 打印模型结构print(model)⚙️ 5. 训练配置importtorch.optimasoptimfromtorch.optimimportlr_scheduler criterionnn.CrossEntropyLoss()# 使用SGD优化器optimizeroptim.SGD(model.parameters(),lr0.001,momentum0.9)# 学习率调度器schedulerlr_scheduler.StepLR(optimizer,step_size7,gamma0.1)️‍♂️ 6. 模型训练循环fromtqdmimporttqdmdeftrain_model(model,dataloaders,criterion,optimizer,scheduler,num_epochs25):best_acc0.0forepochinrange(num_epochs):print(fEpoch{epoch1}/{num_epochs})print(-*10)# 每个epoch有两个阶段训练和验证forphasein[train,val]:ifphasetrain:model.train()dataloaderdataloaders[train]else:model.eval()dataloaderdataloaders[val]running_loss0.0running_corrects0# 进度条withtqdm(dataloader,descphase,leaveFalse)aspbar:forinputs,labelsinpbar:inputsinputs.to(device)labelslabels.to(device)outputsmodel(inputs)losscriterion(outputs,labels)_,predstorch.max(outputs,1)ifphasetrain:optimizer.zero_grad()loss.backward()optimizer.step()running_lossloss.item()*inputs.size(0)running_correctstorch.sum(predslabels.data)ifphasetrain:scheduler.step()epoch_lossrunning_loss/len(dataloaders[phase].dataset)epoch_accrunning_corrects.double()/len(dataloaders[phase].dataset)print(f{phase}Loss:{epoch_loss:.4f}Acc:{epoch_acc:.4f})ifphasevalandepoch_accbest_acc:best_accepoch_acc best_model_wtsmodel.state_dict()print(Training complete)print(fBest Validation Accuracy:{best_acc:.4f})# 加载最佳模型权重model.load_state_dict(best_model_wts)returnmodel# 合并训练和验证的DataLoaderdataloaders{train:train_loader,val:val_loader}# 开始训练modeltrain_model(model,dataloaders,criterion,optimizer,scheduler,num_epochs30)✅ 7. 测试评估defevaluate(model,data_loader,device):model.eval()correct0total0withtorch.no_grad():forinputs,labelsindata_loader:inputsinputs.to(device)labelslabels.to(device)outputsmodel(inputs)_,predictedtorch.max(outputs.data,1)totallabels.size(0)correct(predictedlabels).sum().item()returncorrect/total test_accevaluate(model,test_loader,device)print(fTest Accuracy:{test_acc:.4f}) 8. 可视化预测结果可选importmatplotlib.pyplotaspltimportnumpyasnpdefimshow(inp,titleNone):Imshow for Tensor.inpinp.numpy().transpose((1,2,0))meannp.array([0.485,0.456,0.406])stdnp.array([0.229,0.224,0.225])inpstd*inpmean inpnp.clip(inp,0,1)plt.imshow(inp)iftitle:plt.title(title)plt.pause(0.001)defvisualize_model(model,num_images6):was_trainingmodel.training model.eval()images_so_far0figplt.figure()withtorch.no_grad():fori,(inputs,labels)inenumerate(val_loader):inputsinputs.to(device)labelslabels.to(device)outputsmodel(inputs)_,predstorch.max(outputs,1)forjinrange(inputs.size()[0]):images_so_far1axplt.subplot(num_images//2,2,images_so_far)ax.axis(off)ax.set_title(fPredicted:{val_dataset.classes[preds[j]]})imshow(inputs.cpu().data[j])ifimages_so_farnum_images:model.train(modewas_training)returnmodel.train(modewas_training)visualize_model(model)plt.show()以上文字及代码仅供参考学习使用。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

2026年新手学单片机还要从51开始吗?完整学习路线与避坑指南 2026/9/7 14:22:56

2026年新手学单片机还要从51开始吗?完整学习路线与避坑指南

2026年还有人劝新手学51单片机,是不是有点“老古董”?但我认真告诉你:有,而且应该先学。你去翻大学课程、电子设计竞赛、毕业设计的题目,“51单片机”依然是出现频率最高的关键词之一。尚硅谷今年更新的这套2026新版51…

阅读更多 →
CMSIS-DSP优化原理与嵌入式信号处理工程实践 2026/9/7 14:22:56

CMSIS-DSP优化原理与嵌入式信号处理工程实践

1. 全景拆解:CMSIS-DSP在ARM生态中的定位与模块布局 如果你做过几年嵌入式信号处理相关的固件开发,大概率遇到过这样的场景:项目里要做一个FFT频谱分析,或者要上一套FIR滤波器,手写C代码跑在Cortex-M4上,一…

阅读更多 →
Excel+Word邮件合并批量制作文件夹侧脊标签指南 2026/9/7 14:22:56

Excel+Word邮件合并批量制作文件夹侧脊标签指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
边缘AI实战:ML-KWS-for-MCU源码级解析与TinyML部署指南 2026/9/7 14:22:56

边缘AI实战:ML-KWS-for-MCU源码级解析与TinyML部署指南

1. 项目定位:ML-KWS-for-MCU 为什么值得做源码级审计先说结论:这个仓库是我最近在评估边缘AI落地方案时,翻得最仔细的开源项目之一。ML-KWS-for-MCU(Machine Learning Keyword Spotting for Microcontrollers)是 ARM 维…

阅读更多 →
RDMA技术解析:从零拷贝到内核旁路的分布式系统性能优化实践 2026/9/7 14:22:56

RDMA技术解析:从零拷贝到内核旁路的分布式系统性能优化实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
深挖context_switch:进程切换中mm与内核栈的完整交接机制 2026/9/7 14:19:55

深挖context_switch:进程切换中mm与内核栈的完整交接机制

进程切换这事,表面上看就是调度器挑个新任务,然后“切换上下文”。可真钻进内核代码,你会发现“切换上下文”这四个字背后站着两个完全不同的世界:一个是用户地址空间的搬运,一个是内核栈的接力。我在读 __schedule …

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞