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

从TensorFlow 1 checkpoint迁移Adam参数到PyTorch及state_dict解析

TF1 Adam优化器参数迁移到PyTorch实操指南

一、核心对应关系确认

你推测的TF Adam→PyTorch exp_avg、TF Adam_1→PyTorch exp_avg_sq完全正确,两者分别对应Adam算法中的一阶动量和二阶动量累积值。

二、PyTorch优化器state_dict结构解析

PyTorch的optimizer.state_dict()里,state是一个字典(你看到的[0]/[1]是字典的键,对应模型可训练参数的索引),每个键对应模型中一个requires_grad=True的参数,每个值包含该参数对应的exp_avg、exp_avg_sq(若用AdamW还会有step字段)。

你看到的一维Tensor是PyTorch将每个参数的动量值展平存储的结果,和参数本身的形状一一对应,后续可reshape回原参数形状。

三、参数数量不匹配的原因排查

  1. 可训练参数范围差异:
    • TF1默认仅优化trainable=True的变量,PyTorch优化器默认优化所有requires_grad=True的参数,但注意:
      • PyTorch中nn.BatchNorm/nn.LayerNorm的running_mean、running_var是不可训练的(requires_grad=False),不会被加入优化器,但会被计入模型总参数统计,这是数量差的常见原因。
      • 检查你的TF模型是否有部分变量设置了trainable=False,这些变量对应的Adam动量值不会出现在TF checkpoint中,但你统计PyTorch模型参数时可能包含了对应不可训练参数,导致数量偏差。
  2. 参数分组影响:
    • 如果PyTorch优化器用了参数分组(比如对权重、偏置设不同学习率),state结构本质不变,但需确保分组后的参数顺序和TF可训练参数完全对应。

四、迁移步骤实操

1. 对齐可训练参数列表

先分别导出TF和PyTorch的可训练参数列表,必须保证顺序完全一致:

  • TF1:遍历模型trainable_variables,记录每个变量的名称、形状、总元素数,按顺序排列。
  • PyTorch:遍历model.named_parameters(),筛选p.requires_grad=True的参数,同样记录名称、形状、总元素数,按顺序排列。

2. 从TF checkpoint提取动量值

TF checkpoint中,每个可训练变量var的动量值存储路径通常是:

  • 一阶动量:var.name + '/Adam'
  • 二阶动量:var.name + '/Adam_1'
    提取这些值后,将每个变量的动量值展平为一维数组,按参数顺序拼接成两个大数组(对应所有exp_avg和exp_avg_sq)。

3. 填充PyTorch优化器state_dict

初始化PyTorch优化器(需和原TF Adam超参数一致:lr、betas、eps等),然后遍历填充动量值:

import torch

# 假设已提取TF的一阶动量数组tf_exp_avg、二阶动量数组tf_exp_avg_sq
# 假设已获取PyTorch可训练参数列表pytorch_trainable_params

# 初始化优化器,超参数和TF完全对齐
optimizer = torch.optim.Adam(pytorch_trainable_params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8)

# 遍历参数填充动量值
current_idx = 0
for param in pytorch_trainable_params:
    param_numel = param.numel()
    # 填充一阶动量exp_avg
    optimizer.state[param]['exp_avg'] = torch.tensor(
        tf_exp_avg[current_idx:current_idx+param_numel],
        dtype=param.dtype,
        device=param.device
    ).reshape(param.shape)
    # 填充二阶动量exp_avg_sq
    optimizer.state[param]['exp_avg_sq'] = torch.tensor(
        tf_exp_avg_sq[current_idx:current_idx+param_numel],
        dtype=param.dtype,
        device=param.device
    ).reshape(param.shape)
    current_idx += param_numel

# 保存迁移后的优化器state
torch.save({'optimizer': optimizer.state_dict()}, 'pytorch_optimizer.pt')

4. 验证正确性

  • 检查PyTorch优化器每个参数的exp_avg/exp_avg_sq形状是否和参数本身一致。
  • 运行1步训练,对比损失变化是否和TF1对应步骤一致(允许微小数值精度差异)。

五、特殊情况处理

  • 如果TF checkpoint中动量值是分块存储的,先将分块数组拼接成完整一维数组,再按步骤拆分到对应参数。
  • 如果参数名称不匹配,手动建立TF变量名到PyTorch参数名的映射表,确保参数顺序和数量完全对应。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 18:12:45