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

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/valid padding要对应到PyTorch的具体padding数值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 00:04:58