图像分类模型怎么微调?用 Torchvision 更换分类头并核对验证集

用 Torchvision 的 ResNet18 预训练权重更换二分类头,在独立 train/val 目录训练并输出验证正确数。

已有一小批自有图片,要区分两类对象时,可以从预训练 ResNet18 的分类头开始微调,而不是从零训练。本文用 Torchvision 官方蚂蚁与蜜蜂数据集的目录约定说明输入、训练、验证;代码会输出每轮验证准确率,但两轮演示训练不能证明模型可投入实际业务。

准备互不重叠的 train 与 val

按 PyTorch 官方迁移学习教程准备 data/hymenoptera_data/train/ants、train/bees、val/ants 和 val/bees,各目录放对应图片。教程提供示例数据下载入口;若改用自己的图片,同样保持两类目录名一致,不把同一张或连拍近重复照片分进两边。先手工打开各目录各几张,排除错类或坏图。

图像分类模型怎么微调?用 Torchvision 更换分类头并核对验证集

更换 ResNet18 分类头并运行两轮

安装与你设备匹配的 PyTorch、Torchvision 后,将下列代码保存为 train.py。它沿用官方教程的 ImageFolder、预训练 ResNet18、model.fc 替换和训练/验证分离方法,并把轮数压低用于流程验证。

from pathlib import Path
import torch
from torch import nn, optim
from torch.utils.data import DataLoader
from torchvision import datasets, models, transforms

base = Path("data/hymenoptera_data")
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
])
train = datasets.ImageFolder(base / "train", transform)
val = datasets.ImageFolder(base / "val", transform)
assert train.classes == val.classes, (train.classes, val.classes)
train_loader = DataLoader(train, batch_size=8, shuffle=True)
val_loader = DataLoader(val, batch_size=8, shuffle=False)

device = "cuda" if torch.cuda.is_available() else "cpu"
model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)
model.fc = nn.Linear(model.fc.in_features, len(train.classes))
model = model.to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)

for epoch in range(2):
    model.train()
    for images, labels in train_loader:
        images, labels = images.to(device), labels.to(device)
        optimizer.zero_grad()
        loss = loss_fn(model(images), labels)
        loss.backward()
        optimizer.step()
    model.eval()
    correct = total = 0
    with torch.no_grad():
        for images, labels in val_loader:
            prediction = model(images.to(device)).argmax(dim=1).cpu()
            correct += int((prediction == labels).sum())
            total += len(labels)
    print(f"epoch={epoch + 1}, val_correct={correct}/{total}")

运行 python train.py,预训练权重第一次使用时可能下载。预期每轮出现 val_correct=正确数/验证图片数,具体数字取决于实际图片与环境,本文不伪造。输出只有总准确数;检查实用性还应逐类看错分图片,尤其留意背景、拍摄设备或水印泄漏类别线索。

训练不动或结果虚高怎么办

找不到目录时检查 base 路径;预训练权重下载失败时先确认网络和缓存,不能把随机初始化的模型当成预训练微调;验证集类别目录不一致会触发断言。若验证准确率异常高,查重和近重复照片是否跨集合。真正用于新场景前,需要独立测试集、足够样本和明确错误成本;本例未在用户设备上运行,也不承诺两轮训练有效。

这与手写数字小样本分类演示不同:这里是自有图片目录、Torchvision 预训练权重和独立验证集。官方资料核对于 2026 年 10 月 1 日。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/31578.html

赞 (0)
AI小管家的头像AI小管家
自有图片怎么训练目标检测模型?用 Ultralytics YOLO 检查数据集和验证结果
上一篇 2小时前
用豆包写短视频分镜:把30秒脚本拆成镜头、旁白和拍摄清单
下一篇 2小时前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部