从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回原参数形状。
三、参数数量不匹配的原因排查
- 可训练参数范围差异:
- 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模型参数时可能包含了对应不可训练参数,导致数量偏差。
- PyTorch中
- TF1默认仅优化
- 参数分组影响:
- 如果PyTorch优化器用了参数分组(比如对权重、偏置设不同学习率),
state结构本质不变,但需确保分组后的参数顺序和TF可训练参数完全对应。
- 如果PyTorch优化器用了参数分组(比如对权重、偏置设不同学习率),
四、迁移步骤实操
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
相关产品推荐
相关产品推荐

