实现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()以上文字及代码仅供参考学习使用。

相关新闻

最新新闻

英文Essay降AI使用专业工具和通用大模型哪个更自然

英文Essay降AI使用专业工具和通用大模型哪个更自然

英文Essay降AI使用专业工具和通用大模型哪个更自然 在海外名校商学院与实证金融计量经济学关于面板数据双重差分法(DID)评估碳边境调节机制(CBAM)环境规制效应的课程 Essay 提交过程中,很多留学生都会纠结于工具选型&…

2026/9/7 14:13:36
Tomcat 8.5.100 tar.gz 部署实战:从下载到配置的完整指南

Tomcat 8.5.100 tar.gz 部署实战:从下载到配置的完整指南

简介:Apache Tomcat 8.5.100 是应用广泛的开源 Java 服务器软件,面向需要部署 Servlet 和 JSP 应用的开发者与系统运维人员,可帮助他们快速搭建稳定高效的 Web 运行环境。该版本基于 Java EE 8 规范,支持 Servlet 4.0、JSP 2.3 等…

2026/9/7 14:13:36
三万字大论文整篇降AI应该选择批量处理还是逐段修改

三万字大论文整篇降AI应该选择批量处理还是逐段修改

三万字大论文整篇降AI应该选择批量处理还是逐段修改 在水利水电工程与溃坝水动力学两维浅水波方程数值模拟方向的硕士学位论文送审前夕,很多研究生都会陷入修改策略的选择困境:三万字大论文整篇降AI应该选择批量处理还是逐段修改?整篇长达 3…

2026/9/7 14:13:36
d3 数据变换完全指南:d3.cross、d3.merge、d3.zip 等 8 个数组派生函数详解

d3 数据变换完全指南:d3.cross、d3.merge、d3.zip 等 8 个数组派生函数详解

d3 数据变换完全指南:d3.cross、d3.merge、d3.zip 等 8 个数组派生函数详解 【免费下载链接】d3 Bring data to life with SVG, Canvas and HTML. :bar_chart::chart_with_upwards_trend::tada: 项目地址: https://gitcode.com/GitHub_Trending/d3/d3 D3 的 …

2026/9/7 14:13:36
论文只剩几百字标红时选择按字收费降AI工具更划算吗

论文只剩几百字标红时选择按字收费降AI工具更划算吗

论文只剩几百字标红时选择按字收费降AI工具更划算吗 在法学与知识产权法学专业关于生成式人工智能(AIGC)产出物可版权性认定与侵权责任归责原则方向的硕士学位论文终审修改阶段,很多研究生都会面临收尾阶段的工具选型疑问:论文只…

2026/9/7 14:13:36
像素工厂萌新快速发育攻略:从落地到中期转型不再翻车

像素工厂萌新快速发育攻略:从落地到中期转型不再翻车

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

2026/9/7 14:08:35