PyTorch中如何创建带Dataset类可学习参数的模块并联合训练
问题1:参数收集写法是否正确
这个写法的有效性有一个必要前提:你的自定义Dataset类必须同时继承torch.nn.Module。
只有nn.Module的子类才自带parameters()方法,能自动识别类中用nn.Parameter包裹的可训练参数。如果你只单独继承了torch.utils.data.Dataset,直接调用Dataset.parameters()会直接报错。
你需要按照如下方式定义自定义Dataset,原参数收集代码才是完全正确的:
import torch class TrainableDataset(torch.nn.Module, torch.utils.data.Dataset): def __init__(self): super().__init__() # 必须先初始化nn.Module,才能正常记录可训练参数 # 定义可训练的预处理参数,必须用nn.Parameter包裹 self.learnable_norm_mean = torch.nn.Parameter(torch.tensor([0.5, 0.5, 0.5])) # 其他数据集初始化逻辑... def __getitem__(self, idx): # 数据读取+预处理逻辑...
满足上述定义的前提下,你写的参数收集、优化器初始化代码可以正常工作,优化器会同时更新模型和Dataset内的可训练参数。
问题2:训练/验证数据集的创建逻辑
完全不需要训练完训练集再重新创建验证集,只要保证训练、验证流程共享同一份可训练参数即可,有两种常见实现方案:
- 方案1:拆分参数实例,多数据集共享
你可以把可训练的预处理逻辑单独抽成一个nn.Module子类,创建train_ds和val_ds时都传入同一个参数实例,训练过程中更新的就是这一份共享参数,验证时直接用val_ds就能读取到最新参数的预处理结果:
# 先实例化可训练预处理模块,全程只实例化一次 trainable_preprocess = TrainablePreprocess() # 两个数据集传入同一个预处理模块,共享参数 train_ds = TrainableDataset(preprocess=trainable_preprocess, is_train=True) val_ds = TrainableDataset(preprocess=trainable_preprocess, is_train=False)
- 方案2:单数据集实例,状态切换控制
给自定义Dataset加类似nn.Module的train()/eval()状态切换方法,通过状态变量控制预处理逻辑是否记录梯度,甚至不需要创建两个数据集实例:
class TrainableDataset(torch.nn.Module, torch.utils.data.Dataset): # 其他初始化逻辑不变 def train(self, mode=True): self.training = mode def eval(self): self.train(mode=False) def __getitem__(self, idx): raw_data = self.read_raw_data(idx) if self.training: # 训练模式:预处理计算带梯度,支持参数更新 return self.learnable_preprocess(raw_data) else: # 验证模式:预处理关闭梯度,只用当前参数值做推理 with torch.no_grad(): return self.learnable_preprocess(raw_data)
使用时训练阶段调用ds.train(),验证阶段调用ds.eval()即可。
小提示:如果Dataloader开启了多进程(num_workers>0),需要注意主进程的参数更新要同步到子进程,避免多进程下参数不一致的问题。
内容的提问来源于stack exchange,提问作者Khabbab Zakaria
相关产品推荐
相关产品推荐

