如何基于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.BILINEARresize,这里就不能改;预训练时如果减了均值、除了方差,这里也要同步执行。 - 标签处理:标签必须是整数类型的类别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)]
三、其他训练细节调整
- 学习率:原来的代码里每30个epoch学习率乘以0.1,新数据集可能需要调整这个衰减策略。比如如果新数据集很小,可以用更小的初始学习率(比如原来的1/10),避免模型遗忘预训练的知识。
- 损失函数权重:原来的代码里权重是
[0]+[1 for i in range(17)],如果新数据集的类别分布不平衡,需要调整这个weights_list,比如给样本少的类别更高的权重。 - 验证环节:建议加一个验证步骤,每隔几个epoch在验证集上评估效果,避免过拟合。
内容的提问来源于stack exchange,提问作者Jimbo
相关产品推荐
相关产品推荐

