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

训练采样策略:模型入自定义Dataset还是training_step推理更优?

问题描述

我正在设计训练阶段的样本采样流程,数据格式如下:

defaultdict(list,
        {'C1629836-28004480': [0, 5, 6, 12, 17, 19, 28],
         'C0021846-28004480': [1, 7, 15],
         'C0162832-28004480': [2, 9],
         'C0025929-28004480': [3, 10, 30],
         'C1515655-28004480': [4],
         ...
        }

其中键为标签,值为对应的数据索引列表。我自定义了Dataset类,在其__getitem__(self, idx)函数中需要计算随机选取的锚点与其他数据点的距离,代码片段如下:

def __getitem__(self, idx):
    item_label = self.labels[idx] # C1629836-28004480
    item_data = self.data[item_label] # [0, 5, 6, 12, 17, 19, 28]

    anchor_index = random.sample(item_data,1)
    mentions_indices = [idx for idx in item_data if idx != anchor_index]
    
    with torch.no_grad():
        self.model.eval()
        anchor_input = ...
        anchor_embedding = self.model.mention_encoder(anchor_input)

        for idx in mention_indices: 
        ...

另一种方案是不把模型传入自定义Dataset,转而在训练的training_step函数中执行推理。我了解到Dataset和Dataloader的并行机制能节省训练时间,但这里需要基于训练中模型的最新权重计算距离,不确定并行机制能否保证获取到最新权重(尽管Python变量是引用类型)。实际测试后发现第二种方案速度更快,想请教哪种方式更专业、更正确?

分析与结论

两种方案的核心对比

  • 方案一(Dataset中调用模型):存在明显缺陷

    1. 权重同步问题:Dataloader开启多进程并行时,每个子进程会复制主进程的模型权重副本,主进程更新权重后,子进程的副本不会自动同步,导致计算用的是旧权重,完全不符合你需要「最新权重」的需求。即使单进程模式,把模型耦合进Dataset也违背了数据层与模型层分离的设计原则。
    2. 效率损耗:__getitem__在子进程执行,模型推理会增加进程间通信的额外开销,这也是你测试中方案一更慢的核心原因之一。
  • 方案二(training_step中执行推理):是更专业、合理的选择

    1. 权重一致性:training_step在主进程(或分布式训练的主节点)执行,直接使用当前最新的模型权重,完全不存在权重不同步的问题。
    2. 架构合理性:遵循了数据处理(Dataset/Dataloader负责采样、加载原始数据)与模型计算(training_step负责推理、损失计算)分离的软件工程原则,代码的可维护性、扩展性更强。
    3. 效率优势:你测试中发现的速度更快是合理的——避免了多进程下模型副本的冗余和进程间通信开销,模型推理还能和训练环节整合,借助批量推理、自动混合精度等手段进一步提速。

额外优化建议

如果想进一步提升采样和推理效率,可以:

  • 在Dataset中仅完成锚点与对应样本索引的采样,把原始数据加载好传入training_step,再批量执行模型推理,比逐个样本推理效率更高。
  • 若使用PyTorch Lightning等框架,training_step本身支持批量处理,能更好地利用GPU并行计算能力。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 23:19:53