已关闭
[Bug]: per_token_loss 配置下断点续训崩溃,PrefetchGradAccDataLoader.load_state_dict 位置调用 TypeError #578
iceflysnow创建于 8月7日关闭于 8月20日
8月7日 添加了label:bug
Iiceflysnow
8月7日 关联了pull request:fix(data): accept positional state_dict in PrefetchGradAccDataLoader.load_state_dict
8月7日 关联了pull request:fix(data): accept positional state_dict in PrefetchGradAccDataLoader.load_state_dict
LKONE
8月10日 评论:
8月10日 评论:
您好,反馈的问题经验证确实存在,相关的修改代码PR需要先通过CLA审核,麻烦您先签署一下CLA协议,后续会推动PR合入,感谢您的支持与使用~


iceflysnow
8月10日 评论:
8月10日 评论:
您好,CLA 已签署,CI / mergeable 均正常,请帮忙推动 PR#2949 合入,感谢!


LKONE
8月13日 评论:
8月13日 评论:
相关PR已合入,本issue关闭。
如果后续有其余问题可以再次开放。


8月13日 issue状态由 TODO 改变为 DONE
8月13日 关闭了 issue
8月13日 添加了label:resolved
8月20日 issue状态由 DONE 改变为 ACCEPTED
8月20日 重新打开了 issue
8月20日 issue状态由 ACCEPTED 改变为 DONE
8月20日 关闭了 issue
环境信息
问题描述
当
loss_cfg.loss_type: per_token_loss时,训练侧会用PrefetchGradAccDataLoader包裹基础 dataloader。其load_state_dict签名为 keyword-only(**kwargs),而train_engine.pyload()在 resume 时以位置参数调用它,二者不匹配,断点续训时全 rank 崩溃TypeError,退出码 1,无法从 checkpoint 恢复训练。根因
mindspeed_mm/fsdp/data/dataloader/dataloader.py,PrefetchGradAccDataLoader.load_state_dict为 keyword-only:https://gitcode.com/Ascend/MindSpeed-MM/blob/master/mindspeed_mm/fsdp/data/dataloader/dataloader.py#L252-L253
def load_state_dict(self, **kwargs): # ← keyword-only self.base_dataloader.load_state_dict(**kwargs)mindspeed_mm/fsdp/train/train_engine.pyload()在 resume 时以位置参数调用:https://gitcode.com/Ascend/MindSpeed-MM/blob/master/mindspeed_mm/fsdp/train/train_engine.py#L377-L378
if self.train_dataloader is not None: self.train_dataloader.load_state_dict(state["extra_state"]["train_dataloader"]) # ← 位置传参激活条件(关键)
PrefetchGradAccDataLoader仅在loss_type == "per_token_loss"时启用,否则训练用的是来自torchdata库的基础StatefulDataLoader(其load_state_dict(state_dict)接受位置参数,调用不冲突,故不崩):https://gitcode.com/Ascend/MindSpeed-MM/blob/master/mindspeed_mm/fsdp/train/trainer.py#L381-L385
train_dataloader = build_dataloader(train_dataset) if args.features.loss_cfg.loss_type == "per_token_loss": # ← 仅此 loss_type 才包裹 train_dataloader = PrefetchGradAccDataLoader( train_dataloader, grad_acc_step=args.training.gradient_accumulation_steps )因此本 bug 仅在
per_token_loss配置下浮现(合法且常见的配置,含 Intern-S2 等官方示例)。defaultloss_type 下用基础StatefulDataLoader,位置调用 OK,不崩——这也是为何该 bug 非普遍触发但确属真实缺陷。复现步骤
features: loss_cfg: loss_type: per_token_loss # ← 激活 PrefetchGradAccDataLoader,触发条件 training: lora: enable: true r: 8 save: /path/to/save no_save_optim: false # run1 存全量(含 optim/rng),便于 resume no_save_rng: false no_load_optim: true # run2 从 run1 的 dcp resume no_load_rng: true save_interval: 2 train_iters: 3 load: /path/to/save # run2 resume 用train_iters: 5并指向 run1 的 save 目录续训:预期行为
run2 从 run1 的
iter_0000003干净 resume,iteration 与 consumed_train_samples 连续(如 iter4/5 → iter5/5,consumed 64 → 80),rc=0。实际行为
run2 模型权重成功
Loaded checkpoint后,到 dataloader 恢复处即崩:退出码 1,全 rank 崩。任何
per_token_loss配置的断点续训都会触发。建议修复
让
PrefetchGradAccDataLoader.load_state_dict接受位置参数,与调用方及被包裹的基础StatefulDataLoader(均按位置传参)签名对齐:def state_dict(self): return self.base_dataloader.state_dict() - def load_state_dict(self, **kwargs): - self.base_dataloader.load_state_dict(**kwargs) + def load_state_dict(self, state_dict=None): + # Accept a positional state_dict (the train_engine caller passes it + # positionally). Forward to the wrapped StatefulDataLoader which also + # takes a single positional state_dict argument. + if state_dict is not None: + self.base_dataloader.load_state_dict(state_dict)