You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何基于PyTorch SegNet预训练权重恢复训练并填充批次?

从SegNet预训练权重恢复训练的完整指南

首先,咱们把整个流程拆成两个核心部分:加载预训练权重和填充训练Batch,这刚好对应你问题里的疑惑点。

一、先搞定预训练权重加载

在调用train()函数之前,你得先把已有的预训练权重加载到模型里。假设你的预训练权重文件叫segnet_pretrained.pth,代码可以这么加:

# 初始化模型(和仓库里的写法保持一致)
model = SegNet(input_nbr=input_nbr, label_nbr=label_nbr)
if USE_CUDA:
    model = model.cuda()

# 加载预训练权重
pretrained_weights = torch.load('segnet_pretrained.pth')
model.load_state_dict(pretrained_weights)

# 可选操作:如果新数据集和预训练数据集差异不大,可以冻结前面的卷积层,只训练后面的部分
# for param in model.encoder.parameters():
#     param.requires_grad = False

⚠️ 注意:如果新数据集的类别数和预训练数据集不一样,你需要修改模型最后一层的输出通道数,不然会报错。比如原来预训练是18类(0+17),新数据集是10类,那要重新初始化最后一层的卷积层。

二、核心:填充Batch的代码实现

你贴的代码里# fill the batch的位置,需要完成读取图像、预处理、读取标签、整理成模型所需格式这几个步骤。以下是具体的代码替换,我加了详细注释:

# fill the batch
for idx in range(args.batch_size):
    # 1. 获取当前样本的图像和标签路径(假设batch_files里每个元素是(图像路径, 标签路径)的元组)
    img_path, label_path = batch_files[idx]
    
    # 2. 读取并预处理图像(和预训练时的预处理逻辑必须完全一致!)
    # 示例:用PIL读取图像,转成RGB,resize到指定尺寸
    from PIL import Image
    img = Image.open(img_path).convert('RGB')
    img = img.resize((imsize, imsize), Image.BILINEAR)
    # 转成numpy数组,调整通道顺序为 [通道, 高度, 宽度](PyTorch要求的格式)
    img_np = np.array(img).transpose((2, 0, 1))
    # 归一化(这里用你预训练时用的均值/方差,比如ImageNet的均值或者自己数据集的统计值)
    img_np = img_np / 255.0  # 如果预训练时做了这个归一化步骤,就保留
    # 把处理好的图像放到batch数组对应位置
    batch[idx] = img_np
    
    # 3. 读取并预处理标签
    label = Image.open(label_path).convert('L')  # 假设标签是单通道灰度图,每个像素对应类别ID
    label = label.resize((imsize, imsize), Image.NEAREST)  # 标签必须用最近邻插值,避免出现非整数类别值
    label_np = np.array(label, dtype=int)
    # 把处理好的标签放到batch_labels数组对应位置
    batch_labels[idx] = label_np

关键注意事项:

  • 预处理一致性:图像的resize方法、归一化方式、通道顺序必须和预训练时完全相同,不然模型会出现不兼容的情况。比如预训练时用的是Image.BILINEAR resize,这里就不能改;预训练时如果减了均值、除了方差,这里也要同步执行。
  • 标签处理:标签必须是整数类型的类别ID,绝对不能用线性插值处理标签,否则会出现类似2.5这种无效的类别值。
  • Batch生成:你代码里的batches变量目前是空的,需要提前把新数据集的所有图像-标签路径对分成一个个batch。比如可以用:
    # 假设你已经整理好所有样本的路径列表:all_samples = [(img1, label1), (img2, label2), ...]
    batches = [all_samples[i:i+args.batch_size] for i in range(0, len(all_samples), args.batch_size)]
    

三、其他训练细节调整

  1. 学习率:原来的代码里每30个epoch学习率乘以0.1,新数据集可能需要调整这个衰减策略。比如如果新数据集很小,可以用更小的初始学习率(比如原来的1/10),避免模型遗忘预训练的知识。
  2. 损失函数权重:原来的代码里权重是[0]+[1 for i in range(17)],如果新数据集的类别分布不平衡,需要调整这个weights_list,比如给样本少的类别更高的权重。
  3. 验证环节:建议加一个验证步骤,每隔几个epoch在验证集上评估效果,避免过拟合。

内容的提问来源于stack exchange,提问作者Jimbo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 06:36:41