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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 21:24:03