ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

torchtune 多节点困惑度 PPL 同步计算完整指南

torchtune 多节点困惑度 PPL 同步计算完整指南 torchtune 多节点困惑度 PPL 同步计算完整指南【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune两机 16 卡跑完微调验证阶段各 rank 打印的 loss 却差出一个小数点后第三位——该信哪个或者单机的验证集太大一条 512K 的长上下文直接 OOM只能把数据拆到多台机器上。这就是分布式评估的典型场景torchtune 的多节点同步计算方案让每个 rank 只算自己分片的损失再通过 all_reduce 聚合成与单机一致的全局困惑度PerplexityPPL照下面三步流程照搬即可。 规模上来之后单机评估卡在哪把评估从 1 台机器挪到 N 台机器卡点集中在三处先认全再动手卡点单机的表现多节点的做法数据分片验证集 token 量大单卡串行 forward一次评估要跑数小时用StatefulDistributedSampler把数据集切给 dp 个 rank各算各的跨机聚合只有局部视图每个 rank 只看得见自己的平均 loss对「总损失」和「token 数」各做一次all_reduce(SUM)再求比值精度对齐fp32 累加器跑长循环误差随 token 数漂移聚合缓冲改用 float64最后一步才 exp 还原 PPL关键认知不能直接对各个 rank 的平均 loss求平均。每个 rank 分到的 batch 数、有效 token 数不一样长序列 padding 后差异更大直接平均会引入权重偏差。正确姿势是先把 loss 乘回 token 数聚合成全局的总损失 ÷ 总 token 数。三步完成全局 PPL 聚合分片计算 → all_reduce → 加权平均第 1 步分片计算。每个 rank 在自己的验证分片上 forward用交叉熵损失函数算出 batch 级平均 loss立刻把它乘上该 batch 的有效 token 数即labels ! -100的个数padding 不算得到总损失token 数单独累加。第 2 步all_reduce 求和。对两个累加器各发起一次全局 SUM。数据并行Data ParallelismDP场景下所有 rank 都是全量模型副本这一步没有额外计算量一次 NCCL 集合通信即可。第 3 步加权平均 exp 还原。全局总损失除以全局总 token 数得到加权平均 loss对它取 exp 就是 PPL。整个流程等价于下面 6 行# 1) 本地loss × token 数 → 总损失token 数单独累加 loss_sum float(batch_loss) * n_tokens token_sum n_tokens # 2) 全局各做一次 SUM dist.all_reduce(loss_sum); dist.all_reduce(token_sum) # 3) 加权平均exp 还原 PPL ppl torch.exp(loss_sum / token_sum)注意n_tokens必须排除 paddingloss_sum、token_sum建议建成 float64 缓冲——这是与recipes/full_finetune_distributed.py中validate()实现一致的做法只是把精度提到双精度以压缩多轮累加的漂移。快速跑通从进程组初始化到打印困惑度先初始化进程组。GPU 节点优先用 NCCL 后端跨机带宽利用率显著高于 glootimeout显式设到 180s避免节点多时握手超时直接崩进程import torch.distributed as dist from datetime import timedelta from torchtune.training._distributed import ParallelDims dist.init_process_group(backendnccl, timeouttimedelta(seconds180)) # 2 节点 × 4 卡 8 进程纯数据并行分片 dims ParallelDims(dp_replicate1, dp_shard8, tp1, cp1, world_size8) mesh dims.build_mesh(device_typecuda)ParallelDims在构造时就会校验dp_replicate * dp_shard * tp * cp world_size配错立刻报错而不是跑了一半才挂dp_shard写 -1 可以按world_size // (dp_replicate * tp * cp)自动补齐。纯评估场景把 tp、cp 留 1所有进程都走 dp_shard 即可。再锁种子、建分片加载器。调用 torchtune/training/seed.py 的set_seed(seed42)固定 CPU/CUDA 随机数数据侧用StatefulDistributedSampler(ds, num_replicasdp_degree, rankdp_rank, seed0)num_replicas传 dp 度而不是 world_sizesampler 的 seed 与 recipe 对齐仓库 recipe 中固定为 0保证 8 个 rank 的分片不重不漏。最后跑聚合循环并只在 rank 0 打印model.eval() loss_sum torch.tensor(0.0, dtypetorch.float64, devicedevice) token_sum torch.tensor(0.0, dtypetorch.float64, devicedevice) with torch.no_grad(): for batch in val_loader: n count_valid_tokens(batch) # labels ! -100 的个数 loss_sum float(loss_fn(model(batch), batch)) * n token_sum n dist.all_reduce(loss_sum); dist.all_reduce(token_sum) if dp_rank 0: print(PPL , torch.exp(loss_sum / token_sum).item())fp64 缓冲意味着all_reduce实际是双精度求和token 量到百万级时与单机直算的偏差也能压在 1e-6 量级只让dp_rank 0打印终端就不会刷出 8 行重复结果。若评估对象是用Int4WeightOnlyQuantizer(groupsize128)量化的模型PPL 天然高于原始精度权重先在同一份量化配置下单机跑一遍基线再判断多节点聚合本身有没有问题。 结果对不上时的排查清单如果出现各 rank 局部指标都正常、唯独全局 PPL 偏大先检查聚合缓冲的 dtype——fp32 在长循环里误差随 token 数累积把loss_sum/token_sum改成 float64 重跑对比。如果出现两机 PPL 与单机差值大于 1e-3先检查分片一致性num_replicas是否误传了 world_size、sampler seed 是否 0、有没有人绕过 sampler 手写了dataset[i::k]式切片。如果出现某个 rank 报 NCCL watchdog 超时退出先检查init_process_group的timeout建议 ≥180s并确认所有节点走同一后端——一端 nccl、一端 gloo 必然握手卡死。如果出现多次运行结果漂移先检查种子开跑前所有 rank 是否执行过set_seedresume 场景下 checkpoint 里记录的 seed 是否与配置一致recipe 恢复时会校验并报警。源码里看这几处就够模块路径一句话职责torchtune/training/_distributed.pyParallelDims维度校验、mesh 构建与分布式初始化入口recipes/full_finetune_distributed.pyvalidate()是 all_reduce 聚合损失与 token 数、算加权平均的参考实现torchtune/training/seed.pyset_seed固定全局随机数评估可复现性的基础torchtune/training/quantization.pyInt4WeightOnlyQuantizer等量化器低精度评估入口把分片、聚合、exp 三步串起来后多机与单机的口径就统一了剩下的只是选对后端和超时参数。常用入口如下docs/source/overview.rsttorchtune 项目总览与配置体系说明recipes/configs/llama3_1/evaluation.yaml评估任务的现成配置参考torchtune/training/分布式、种子、量化等训练工具集tests/recipes/test_full_finetune_distributed.py分布式 recipe 的回归测试聚合逻辑改动时可先跑它【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表