在Python中如何将Keras训练好的模型权重转换为PyTorch格式?

千枫酱_1332

千枫酱_1332

2026-06-15

412人浏览

原创

不能直接加载keras的.h5权重到pytorch,因二者模型结构、张量布局(channels_last vs channels_first)、参数命名(gamma/weight、moving_mean/running_mean)及权重形状(conv2d为kh×kw×in×out→out×in×kh×kw)均不同,需手动逐层对齐、转置、重命名并校验数值一致性。

在python中如何将keras训练好的模型权重转换为pytorch格式?

为什么不能直接加载Keras的.h5权重到PyTorch

Keras(尤其是TensorFlow后端)和PyTorch的模型结构、张量布局、参数命名规则完全不同。比如:Conv2D在Keras中默认data_format='channels_last',而PyTorch是channels_first;Dense层的权重矩阵在Keras里是(in, out),PyTorch里是(out, in);还有BN层的gamma/beta vs weight/bias、moving_mean vs running_mean等映射差异。直接用torch.load()读.h5会报错或加载错位。

手动逐层映射权重的实操步骤

核心是:先用Keras读权重,再按层名/类型对齐PyTorch模型结构,做形状转置+名称转换+数值拷贝。需确保两个模型拓扑完全一致(层数、类型、通道数、kernel size等)。

  • 用h5py.File('model.h5', 'r')打开权重文件,遍历model.layers或file.keys()定位每层参数路径,例如conv2d_1/kernel:0对应卷积核,batch_normalization_1/gamma:0对应缩放参数
  • PyTorch模型中对应层需提前创建好(不能靠load_state_dict()自动匹配),用named_parameters()和named_buffers()确认目标键名,如features.0.weight、features.1.running_mean
  • 对卷积层:Keras权重shape为(kh, kw, in_c, out_c) → 转为(out_c, in_c, kh, kw)再赋给torch_tensor.data.copy_()
  • 对BN层:Keras的gamma→PyTorch的weight,beta→bias,moving_mean→running_mean,moving_variance→running_var;注意running_var需开方后再取倒数(因PyTorch BN内部用1/sqrt(var + eps))

用keras2pytorch这类工具库的风险点

社区有keras2pytorch或tf2pytorch等小众转换脚本,但它们只覆盖常见层(Conv2D、Dense、BatchNormalization),遇到DepthwiseConv2D、自定义Lambda层、RNN或tf.keras.layers.Resizing就大概率失败。更麻烦的是——这些工具不校验数值一致性。

python-pro
python-pro

高级 Python 特性、异步编程、性能调优、静态类型、内存管理、Python 内部机制及生态库方面的专家。

下载
  • 务必在转换后用同一组输入跑前向,比对Keras输出和PyTorch输出的torch.allclose(out_keras, out_pt, atol=1e-5)
  • 如果Keras模型用了tf.keras.applications里的预训练模型,优先查官方是否提供PyTorch版(如torchvision.models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V1)),避免自己转换
  • 若原Keras模型含tf.keras.layers.GlobalAveragePooling2D,PyTorch对应是nn.AdaptiveAvgPool2d((1,1)),但权重无参数,只需保证输出shape一致即可

保存为PyTorch可复用的格式

转换完别只存state_dict,否则下次加载还得重建模型结构。推荐两种方式:

  • 保存完整模型:torch.save(model, 'model.pt')(要求模型类定义在运行时可见,且无依赖外部闭包)
  • 更稳妥的做法:保存state_dict + 单独存一份模型结构代码(如model_def.py),加载时先import类再model = MyModel(); model.load_state_dict(torch.load('weights.pth'))
  • 避免用torch.save(model.state_dict(), ...)后忘记记录输入尺寸/归一化参数——PyTorch没Keras的model.input_shape元信息,这些必须手动维护

权重转换不是格式搬运,本质是跨框架的数值重布线。最易被忽略的是BN层的training=False模式下running_var的数值处理,以及全局池化层前后张量permute逻辑是否隐式存在。动手前先用最小网络(单卷积+BN)跑通全流程,再放大到完整模型。

Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!

相关文章

PHP速学视频免费教程(入门到精通)
PHP速学视频免费教程(入门到精通)

PHP怎么学习?PHP怎么入门?PHP在哪学?PHP怎么学才快?不用担心,这里为大家提供了PHP速学教程(入门到精通),有需要的小伙伴保存下载就能学习啦!

下载

相关标签:

pytorch python

本站声明:本文内容由网友自发贡献,版权归原作者所有,本站不承担相应法律责任。如您发现有涉嫌抄袭侵权的内容,请联系admin@php.cn

相关专题

更多
python打包成可执行文件
python打包成可执行文件

本专题为大家带来python打包成可执行文件相关的文章,大家可以免费的下载体验。

2023.07.20

1671

4

python能做什么
python能做什么

python能做的有:可用于开发基于控制台的应用程序、多媒体部分开发、用于开发基于Web的应用程序、使用python处理数据、系统编程等等。本专题为大家提供python相关的各种文章、以及下载和课程。

2023.07.25

4224

7

format在python中的用法
format在python中的用法

Python中的format是一种字符串格式化方法,用于将变量或值插入到字符串中的占位符位置。通过format方法,我们可以动态地构建字符串,使其包含不同值。php中文网给大家带来了相关的教程以及文章,欢迎大家前来阅读学习。

2023.07.31

1669

3

python教程
python教程

Python已成为一门网红语言,即使是在非编程开发者当中,也掀起了一股学习的热潮。本专题为大家带来python教程的相关文章,大家可以免费体验学习。

2023.08.03

24537

23

python环境变量的配置
python环境变量的配置

Python是一种流行的编程语言,被广泛用于软件开发、数据分析和科学计算等领域。在安装Python之后,我们需要配置环境变量,以便在任何位置都能够访问Python的可执行文件。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

3007

5

python eval
python eval

eval函数是Python中一个非常强大的函数,它可以将字符串作为Python代码进行执行,实现动态编程的效果。然而,由于其潜在的安全风险和性能问题,需要谨慎使用。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.04

3027

5

scratch和python区别
scratch和python区别

scratch和python的区别:1、scratch是一种专为初学者设计的图形化编程语言,python是一种文本编程语言;2、scratch使用的是基于积木的编程语法,python采用更加传统的文本编程语法等等。本专题为大家提供scratch和python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

1163

5

python合并两个列表
python合并两个列表

Python是一种强大的编程语言,具有许多方便的功能和工具。在Python中,有多种方法可以合并两个列表。php中文网给大家带来了相关的教程以及文章,欢迎大家前来学习阅读。

2023.08.10

596

4

python是前端还是后端
python是前端还是后端

Python属于前端也属于后端,其灵活性和丰富的生态系统使得开发人员能够在不同的领域中灵活运用。本专题为大家提供python相关的文章、下载、课程内容,供大家免费下载体验。

2023.08.11

2343

5

热门下载

更多
网站特效
/
网站源码
/
网站素材
/
前端模板

精品课程

更多
相关推荐
/
热门推荐
/
最新课程