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

使用Huggingface Transformers Trainer从Checkpoint恢复训练时遇设备错误

解决BERT恢复训练时的设备不匹配及fused参数错误

核心问题分析

  • 恢复训练时,仅把模型移到CUDA不够,优化器、学习率调度器的状态张量可能还留在CPU,导致设备不匹配。
  • 升级Transformers后出现的fused=True错误,是因为fused优化器要求所有参数必须是CUDA浮点张量,而从checkpoint加载的部分参数不符合要求。

分步解决方案

方案1:完整加载checkpoint并统一设备

  1. 加载模型时直接指定设备映射,确保参数直接加载到CUDA:
model = BertForMaskedLM.from_pretrained(
    'mybert\checkpoint-5410440',
    device_map="cuda:0"  # 直接映射到目标CUDA设备
)
  1. 恢复训练时指定checkpoint路径,而非仅设resume_from_checkpoint=True:
trainer.train(resume_from_checkpoint='mybert\checkpoint-5410440')

注:部分版本中resume_from_checkpoint=True存在设备映射bug,指定具体路径更可靠。

方案2:手动迁移优化器/调度器状态到CUDA

如果已经初始化了Trainer,恢复前手动将优化器和调度器的状态张量移到CUDA:

# 迁移优化器状态
for state in trainer.optimizer.state.values():
    for k, v in state.items():
        if isinstance(v, torch.Tensor):
            state[k] = v.to("cuda:0")

# 迁移学习率调度器状态
if trainer.lr_scheduler is not None:
    for key, val in trainer.lr_scheduler.state_dict().items():
        if isinstance(val, torch.Tensor):
            trainer.lr_scheduler.state_dict()[key] = val.to("cuda:0")

方案3:解决fused=True参数错误

升级后出现的fused错误,可通过两种方式处理:

  • 禁用fused优化器:在TrainingArguments中改用非fused版本的AdamW:
training_args = TrainingArguments(
    # 其他参数...
    optim="adamw_torch"  # 替换默认的fused版本
)
  • 若坚持用fused优化器,强制将所有模型参数转为CUDA浮点张量:
# 遍历模型参数,统一设备和数据类型
for param in model.parameters():
    param.data = param.data.to(torch.float32).to("cuda:0")
    if param.grad is not None:
        param.grad.data = param.grad.data.to(torch.float32).to("cuda:0")

额外注意事项

  • 检查数据加载器:确保输入数据在加载时也移到CUDA,避免数据在CPU上导致设备不匹配。
  • 清理旧缓存:如果之前的checkpoint存在混合设备状态,可删除后先跑一小步生成新checkpoint,再尝试恢复。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 15:22:15