TensorFlow转PyTorch:Unet权重迁移后输出不一致问题排查
我有一个基于TensorFlow实现的Unet模型,已在自定义数据集上完成训练,权重保存为.hdf5格式。现需将代码迁移至PyTorch,且已实现结构等效的PyTorch模型,但在迁移权重时遇到问题。我通过层间复制的方式将TensorFlow权重转换为PyTorch的state_dict(代码如下),但加载权重后的PyTorch模型输出与原TensorFlow模型差异极大(输出混乱)。
我怀疑问题出在权重转置环节,但不知如何修复,恳请相关排查思路与解决指导。
def weight_loading(pretrained_weights): # Load the weights tf_model = tf.keras.models.load_model(pretrained_weights) tf_weights = tf_model.get_weights() # Load the PyTorch model pt_model = UNet() #implemented based on the previous model (by myself) initial_state_dict = pt_model.state_dict() new_state_dict = {} with torch.no_grad(): x = 0 for i, layer in enumerate(pt_model.modules()): if isinstance(layer, torch.nn.Conv2d): # extract the weights and biases from the TensorFlow weights weight_tf = tf_weights[x*2] bias_tf = tf_weights[x*2+1] # convert the weights and biases to PyTorch format weight_pt = torch.tensor(weight_tf.transpose()) bias_pt = torch.tensor(bias_tf) # get the name of the weight and bias tensors weight_name = list(pt_model.named_parameters())[x*2][0] bias_name = list(pt_model.named_parameters())[x*2+1][0] # set the weights and biases in the PyTorch model state_dict new_state_dict[weight_name]= weight_pt new_state_dict[bias_name] = bias_pt x = x + 1 if isinstance(layer, torch.nn.ConvTranspose2d): weight_tf = tf_weights[x*2] bias_tf = tf_weights[x*2+1] # convert the weights and biases to PyTorch format weight_pt = torch.tensor(np.transpose(weight_tf, (2, 3, 0, 1))) bias_pt = torch.tensor(bias_tf) # get the name of the weight and bias tensors weight_name = list(pt_model.named_parameters())[x*2][0] bias_name = list(pt_model.named_parameters())[x*2+1][0] # set the weights and biases in the PyTorch model state_dict new_state_dict[weight_name] = weight_pt new_state_dict[bias_name] = bias_pt x = x + 1 # load the new generated state_dict to pt_model pt_model.load_state_dict(new_state_dict) return pt_model
注:我已按层复制权重,涉及Conv2d与ConvTranspose2d层,期望加载权重后的PyTorch模型与原TensorFlow模型输出一致,但实际差异极大。更新:检查发现第一次最大池化(两层卷积后)的输出略有相似,但后续差异明显。
1. 核心错误:权重维度转置逻辑错误
TensorFlow和PyTorch的卷积层权重维度定义完全不同,你的转置逻辑存在根本性错误:
TensorFlow Conv2D权重维度:
(kernel_height, kernel_width, in_channels, out_channels)PyTorch Conv2d权重维度:
(out_channels, in_channels, kernel_height, kernel_width)
你当前用weight_tf.transpose()完全错误,正确转置应为np.transpose(weight_tf, (3, 2, 0, 1))TensorFlow Conv2DTranspose权重维度:
(kernel_height, kernel_width, out_channels, in_channels)(注意和正向卷积的in/out顺序相反)PyTorch ConvTranspose2d权重维度:
(in_channels, out_channels, kernel_height, kernel_width)
你当前的转置(2,3,0,1)错误,正确转置应为np.transpose(weight_tf, (3, 2, 0, 1))
2. 层顺序匹配问题
你用pt_model.modules()遍历层会包含所有子模块(比如Sequential里的子层),同时用list(pt_model.named_parameters())[x*2]匹配参数名,极易出现层顺序不对应的问题——TensorFlow的权重列表是按模型定义的层顺序排列,而PyTorch的modules遍历、named_parameters顺序可能和TF不匹配。
修复方式:
- 分别打印TensorFlow模型各Conv/ConvTranspose层的名称、权重形状
- 打印PyTorch模型
named_parameters()的名称、形状,手动建立两者的层权重映射关系,再按对应关系复制权重,不要依赖自动遍历顺序。
3. 数据格式与预处理差异
- TensorFlow默认输入格式是
(batch_size, height, width, channels),PyTorch是(batch_size, channels, height, width),如果输入未转置通道维度,会直接导致输出完全错误。 - 检查归一化/标准化逻辑是否完全一致:比如原TF模型的输入归一化均值、方差,是否和PyTorch模型一致,是否存在数据范围差异(比如TF用[0,1],PyTorch用[-1,1])。
4. 中间输出定位错误层
既然第一次池化后输出还有相似性,说明前几层权重可能部分正确,后续层出错。可以:
- 分别提取TF和PyTorch模型每一层的输出(第一个Conv、第二个Conv、第一个池化、第一个上采样等),逐一对比数值差异,定位到具体哪一层开始出现明显偏差,再针对性检查该层的权重转置和层匹配是否正确。
修复后的权重加载代码示例
修正转置逻辑,同时确保层顺序匹配(需根据你的模型结构调整映射关系):
import numpy as np import tensorflow as tf import torch def weight_loading(pretrained_weights): # Load TF model and weights tf_model = tf.keras.models.load_model(pretrained_weights) tf_weights = tf_model.get_weights() # Load PyTorch model pt_model = UNet() pt_state_dict = pt_model.state_dict() # 先手动确认TF和PyTorch的层顺序,再分配索引 tf_conv_count = len([l for l in tf_model.layers if isinstance(l, tf.keras.layers.Conv2D)]) tf_conv_idx = 0 tf_deconv_idx = tf_conv_count * 2 # 每个Conv对应权重+偏置,共2*count个元素 with torch.no_grad(): # 处理Conv2d层 for name, param in pt_state_dict.items(): if 'conv' in name and 'weight' in name and 'transpose' not in name: # TF Conv2D权重:(kh, kw, in_ch, out_ch) tf_weight = tf_weights[tf_conv_idx] # 转置为PyTorch格式:(out_ch, in_ch, kh, kw) pt_weight = torch.tensor(np.transpose(tf_weight, (3, 2, 0, 1)), dtype=torch.float32) param.copy_(pt_weight) tf_conv_idx += 1 # 处理对应偏置 tf_bias = tf_weights[tf_conv_idx] pt_bias = torch.tensor(tf_bias, dtype=torch.float32) pt_state_dict[name.replace('weight', 'bias')].copy_(pt_bias) tf_conv_idx += 1 # 处理ConvTranspose2d层 for name, param in pt_state_dict.items(): if 'transpose' in name and 'weight' in name: # TF Conv2DTranspose权重:(kh, kw, out_ch, in_ch) tf_weight = tf_weights[tf_deconv_idx] # 转置为PyTorch格式:(in_ch, out_ch, kh, kw) pt_weight = torch.tensor(np.transpose(tf_weight, (3, 2, 0, 1)), dtype=torch.float32) param.copy_(pt_weight) tf_deconv_idx += 1 # 处理对应偏置 tf_bias = tf_weights[tf_deconv_idx] pt_bias = torch.tensor(tf_bias, dtype=torch.float32) pt_state_dict[name.replace('weight', 'bias')].copy_(pt_bias) tf_deconv_idx += 1 pt_model.load_state_dict(pt_state_dict) return pt_model
额外注意事项
- 检查激活函数:比如TF用
relu6的话,PyTorch要实现相同的截断逻辑,不能直接用默认ReLU。 - 检查池化层:TF的
MaxPool2D和PyTorch的MaxPool2d的kernel_size、stride、padding参数要完全一致,特别是TF的same/validpadding要对应到PyTorch的具体padding数值。
内容的提问来源于stack exchange,提问作者farzaneh

