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

使用PyTorch的DistributedSampler做推理时如何匹配样本与对应预测结果

分布式推理样本对齐最简方案

核心思路是给每个样本绑定全局原始索引,随数据一同进入推理流程,无需解码即可直接对齐原始样本。

步骤1:改造数据集结构

在构造TensorDataset时,额外加入和原始样本顺序完全对应的全局索引张量:

# 生成和原始输入顺序一一对应的全局索引
sample_global_idx = torch.arange(len(encoded_dict['input_ids']), dtype=torch.long)
# 将索引加入数据集,和原有张量按顺序对应
dataset = TensorDataset(encoded_dict['input_ids'], encoded_dict['attention_mask'], sample_global_idx)
sampler = DistributedSampler(
    dataset, num_replicas=args.nodes * args.gpus, rank=args.node_rank * args.gpus + gpu_number, shuffle=False
)
dataloader = DataLoader(dataset, batch_size=batch_size, sampler=sampler)

步骤2:关联索引与预测结果

推理迭代dataloader时,取出每个batch对应的索引,将预测结果与索引成对保存:

for batch in dataloader:
    input_ids, attn_mask, batch_idx = batch
    # 正常前向推理得到预测结果
    preds = model(input_ids, attention_mask=attn_mask)
    # 保存时绑定索引与对应预测结果,可存在本地数组、临时文件等
    saved_results.extend(zip(batch_idx.tolist(), preds.tolist()))

步骤3:合并排序得到对齐结果

所有进程推理完成后,收集所有进程保存的(索引,预测结果)对,按索引从小到大排序,得到的结果顺序和原始输入样本顺序完全一致。

注意:如果样本总数无法被总GPU数整除,DistributedSampler默认会复制末尾样本填充到可整除长度,收集结果后过滤掉索引超过原始样本最大序号的重复项即可。如果需要避免填充,可在初始化DistributedSampler时设置drop_last=True,直接丢弃末尾不足一个分块的样本。

如果希望省去多进程文件合并的步骤,也可以直接调用torch.distributed.gather接口,把所有进程的结果收集到主进程后再排序。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 06:54:02