实现24类花卉的高精度分类 PyTorch训练花卉分类数据集24类 使用花卉数据集进行图像分类
2026/9/7 13:43:32 网站建设 项目流程

采用预训练模型(如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# 数据增强和标准化transform=transforms.Compose([transforms.Resize((224,224)),transforms.ToTensor(),transforms.Normalize(mean=[0.485,0.456,0.406],std=[0.229,0.224,0.225])])# 数据集路径data_dir='flowers_dataset'train_dataset=datasets.ImageFolder(os.path.join(data_dir,'train'),transform=transform)val_dataset=datasets.ImageFolder(os.path.join(data_dir,'val'),transform=transform)test_dataset=datasets.ImageFolder(os.path.join(data_dir,'test'),transform=transform)# DataLoaderbatch_size=64train_loader=DataLoader(train_dataset,batch_size=batch_size,shuffle=True,num_workers=4)val_loader=DataLoader(val_dataset,batch_size=batch_size,shuffle=False,num_workers=4)test_loader=DataLoader(test_dataset,batch_size=batch_size,shuffle=False,num_workers=4)print("Number of classes:",len(train_dataset.classes))print("Class names:",train_dataset.classes)

🤖 4. 模型定义(使用 ResNet50)

importtorchimporttorch.nnasnnfromtorchvisionimportmodels device=torch.device("cuda"iftorch.cuda.is_available()else"cpu")# 使用预训练的ResNet50model=models.resnet50(pretrained=True)# 修改最后一层全连接层,适配24类num_ftrs=model.fc.in_features model.fc=nn.Linear(num_ftrs,24)# 24种花卉model=model.to(device)# 打印模型结构print(model)

⚙️ 5. 训练配置

importtorch.optimasoptimfromtorch.optimimportlr_scheduler criterion=nn.CrossEntropyLoss()# 使用SGD优化器optimizer=optim.SGD(model.parameters(),lr=0.001,momentum=0.9)# 学习率调度器scheduler=lr_scheduler.StepLR(optimizer,step_size=7,gamma=0.1)

🏋️‍♂️ 6. 模型训练循环

fromtqdmimporttqdmdeftrain_model(model,dataloaders,criterion,optimizer,scheduler,num_epochs=25):best_acc=0.0forepochinrange(num_epochs):print(f'Epoch{epoch+1}/{num_epochs}')print('-'*10)# 每个epoch有两个阶段:训练和验证forphasein['train','val']:ifphase=='train':model.train()dataloader=dataloaders['train']else:model.eval()dataloader=dataloaders['val']running_loss=0.0running_corrects=0# 进度条withtqdm(dataloader,desc=phase,leave=False)aspbar:forinputs,labelsinpbar:inputs=inputs.to(device)labels=labels.to(device)outputs=model(inputs)loss=criterion(outputs,labels)_,preds=torch.max(outputs,1)ifphase=='train':optimizer.zero_grad()loss.backward()optimizer.step()running_loss+=loss.item()*inputs.size(0)running_corrects+=torch.sum(preds==labels.data)ifphase=='train':scheduler.step()epoch_loss=running_loss/len(dataloaders[phase].dataset)epoch_acc=running_corrects.double()/len(dataloaders[phase].dataset)print(f'{phase}Loss:{epoch_loss:.4f}Acc:{epoch_acc:.4f}')ifphase=='val'andepoch_acc>best_acc:best_acc=epoch_acc best_model_wts=model.state_dict()print('Training complete')print(f'Best Validation Accuracy:{best_acc:.4f}')# 加载最佳模型权重model.load_state_dict(best_model_wts)returnmodel# 合并训练和验证的DataLoaderdataloaders={'train':train_loader,'val':val_loader}# 开始训练model=train_model(model,dataloaders,criterion,optimizer,scheduler,num_epochs=30)

✅ 7. 测试评估

defevaluate(model,data_loader,device):model.eval()correct=0total=0withtorch.no_grad():forinputs,labelsindata_loader:inputs=inputs.to(device)labels=labels.to(device)outputs=model(inputs)_,predicted=torch.max(outputs.data,1)total+=labels.size(0)correct+=(predicted==labels).sum().item()returncorrect/total test_acc=evaluate(model,test_loader,device)print(f'Test Accuracy:{test_acc:.4f}')

📊 8. 可视化预测结果(可选)

importmatplotlib.pyplotaspltimportnumpyasnpdefimshow(inp,title=None):"""Imshow for Tensor."""inp=inp.numpy().transpose((1,2,0))mean=np.array([0.485,0.456,0.406])std=np.array([0.229,0.224,0.225])inp=std*inp+mean inp=np.clip(inp,0,1)plt.imshow(inp)iftitle:plt.title(title)plt.pause(0.001)defvisualize_model(model,num_images=6):was_training=model.training model.eval()images_so_far=0fig=plt.figure()withtorch.no_grad():fori,(inputs,labels)inenumerate(val_loader):inputs=inputs.to(device)labels=labels.to(device)outputs=model(inputs)_,preds=torch.max(outputs,1)forjinrange(inputs.size()[0]):images_so_far+=1ax=plt.subplot(num_images//2,2,images_so_far)ax.axis('off')ax.set_title(f'Predicted:{val_dataset.classes[preds[j]]}')imshow(inputs.cpu().data[j])ifimages_so_far==num_images:model.train(mode=was_training)returnmodel.train(mode=was_training)visualize_model(model)plt.show()

以上文字及代码仅供参考学习使用。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询