从ResNet18到ONNX再到.wts:一站式模型转换与权重导出指南
1. 为什么你需要掌握模型转换这条“流水线”?
如果你玩过AI模型,尤其是用PyTorch或者TensorFlow训练过自己的网络,那你肯定对.pth、.ckpt这类文件不陌生。它们就像是模型的“家”,里面既有房子的设计图纸(模型结构),也装满了家具(模型权重)。在自己熟悉的框架环境里,用起来当然很顺手。但问题来了,当你辛辛苦苦训练好一个模型,想把它部署到其他地方——比如一个没有PyTorch环境的嵌入式设备、一个追求极致推理速度的服务器,或者一个特定的推理引擎里时,你可能会发现,这个“家”搬不过去。
这时候,模型格式转换就成了必须掌握的技能。这就像你要把一套家具从宜家(PyTorch)搬到另一个国家,你得先把它们拆解成标准化的零件(ONNX),再根据新家的安装说明书(如TensorRT、OpenVINO等)重新组装,甚至打包成更轻便的运输箱(.wts)。今天,我就以最经典的图像分类模型ResNet18为例,带你走通这条从PyTorch到ONNX,再到.wts权重文件的完整转换流水线。我保证,就算你之前没接触过模型转换,跟着我的步骤走,也能轻松搞定。
我选择ResNet18,不仅因为它结构经典、资料丰富,更因为它在实际部署中非常常见,从智能摄像头到工业质检,都能看到它的身影。而ONNX(Open Neural Network Exchange)格式,是目前业界最通用的模型交换“中间语”,它能让你的模型在不同框架和硬件之间自由穿梭。最后的.wts文件,则是像NVIDIA TensorRT这类高性能推理引擎偏爱的权重存储格式之一,特别适合在资源受限的环境下加载。
所以,无论你是想把自己的模型塞进Jetson Nano这样的小型设备,还是为了提升线上服务的推理性能,这条转换路径都是你的必修课。别担心复杂,我踩过的坑、总结的技巧,都会在这篇文章里分享给你,咱们一步步来。
2. 第一步:获取并保存你的ResNet18模型
万事开头难?在这里一点也不。获取一个预训练的ResNet18模型,在PyTorch里可以说是最简单的事情之一。但“保存”这个动作,却有几个关键的细节决定了你后续转换的成败。
首先,我们来看看最直接的代码。下面的脚本会下载PyTorch官方预训练的ResNet18模型,并把它保存为resnet18.pth文件。我强烈建议你新建一个Python文件,比如叫download_resnet18.py,把代码复制进去运行。
import torch
import torchvision
def main():
# 检查CUDA是否可用,这会影响模型加载的设备
print('CUDA device count: ', torch.cuda.device_count())
# 下载预训练的ResNet18模型
# 注意:`pretrained=True` 参数在torchvision新版本中可能已变更,但当前主流版本仍支持
net = torchvision.models.resnet18(pretrained=True)
# 将模型转移到GPU(如果可用)并设置为评估模式
device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
net = net.to(device)
net.eval() # 这一步至关重要!它关闭了Dropout、BatchNorm的随机性
# 打印模型结构,确认加载成功
print(net)
# 生成一个随机输入,测试模型前向传播是否正常
# ResNet18的标准输入尺寸是3通道,224x224像素
dummy_input = torch.ones(1, 3, 224, 224).to(device)
output = net(dummy_input)
print('ResNet18 output shape:', output.shape) # 应该是 torch.Size([1, 1000])
# 保存整个模型
torch.save(net, "resnet18.pth")
print("Model saved as 'resnet18.pth'")
if __name__ == '__main__':
main()
运行这段代码,你会得到一个resnet18.pth文件。这个文件是PyTorch默认的保存格式,它采用Python的pickle机制,同时序列化了模型的结构定义和所有的权重参数。这种保存方式非常方便,因为加载时只需要一行torch.load,模型结构和权重就都回来了。
但是,这里有几个我踩过坑的要点要提醒你:
.eval()模式:在保存模型之前,务必调用net.eval()。这行代码会让模型中的BatchNorm层和Dropout层固定下来,使用训练好的统计参数,而不是在推理时继续变化。如果忘记这一步,在转换ONNX或推理时可能会得到不一致的结果。- 设备一致性:我们通常会在GPU上训练和验证模型,但保存时,PyTorch的
torch.save会连同模型所在的设备信息一起保存。如果你在GPU上保存,然后在没有GPU的机器上加载,可能会报错。一个更稳健的做法是在保存前将模型转移到CPU:net = net.to('cpu'),然后再保存。不过对于后续的ONNX转换,我们通常会在有GPU的环境下进行,所以这里按原样保存问题不大。 - 输入尺寸:我们用了
(1, 3, 224, 224)作为测试输入。记住这个尺寸,因为在转换ONNX时,我们需要提供一个同样尺寸的“样例输入”,ONNX模型会记录这个输入的形状。
保存好.pth文件,我们的原材料就准备好了。接下来,我们要把它转换成更通用的格式。
3. 第二步:将PyTorch模型转换为ONNX格式
拿到resnet18.pth后,我们就可以进行第一次“转码”了:把它变成ONNX格式。你可以把ONNX想象成一个“通用翻译器”。PyTorch、TensorFlow、MXNet等框架说的话(模型格式)各不相同,但它们都能把自己的模型翻译成ONNX这种“世界语”。然后,任何支持ONNX的推理引擎(如TensorRT, OpenVINO, ONNX Runtime)都能读懂并执行它。
转换的核心是PyTorch内置的torch.onnx.export函数。这个函数功能强大,但参数也不少。下面是我常用的转换脚本,同样,建议你保存为convert_to_onnx.py。
import torch
def main():
# 1. 加载之前保存的模型
# 注意:这里加载的是包含结构和权重的完整模型文件
model = torch.load('resnet18.pth', map_location='cpu') # 先加载到CPU上更稳妥
# 2. 将模型设置为评估模式(再次确认)
model.eval()
# 3. 创建一个虚拟输入(dummy input)
# 这个输入的维度必须和模型期望的完全一致,batch size=1, channels=3, height=224, width=224
# 我们通常在CPU上创建这个输入,因为export过程不强制需要GPU
dummy_input = torch.randn(1, 3, 224, 224)
# 4. 指定输出ONNX文件的路径
onnx_file_path = 'resnet18.onnx'
# 5. 执行导出
torch.onnx.export(
model, # 要转换的模型
dummy_input, # 模型输入样例
onnx_file_path, # 输出文件路径
input_names=['input'], # 输入节点的名称
output_names=['output'], # 输出节点的名称
opset_version=11, # ONNX算子集版本,11是一个广泛支持的稳定版本
dynamic_axes={'input': {0: 'batch_size'}, # 指定动态维度,这里让batch size可变
'output': {0: 'batch_size'}}
)
print(f"Model has been converted to ONNX format and saved as '{onnx_file_path}'")
if __name__ == '__main__':
main()
运行这个脚本,你就会得到resnet18.onnx文件。这个文件是二进制的,你可以用Netron(一个超好用的模型可视化工具)打开它,直观地看到整个ResNet18的计算图结构,每一层叫什么,输入输出是什么,清清楚楚。
这里我着重解释两个容易出问题的地方:
opset_version:这是ONNX的算子集版本号。ONNX在不断更新,新的版本会支持更多、更高效的算子。版本号太低,可能不支持你模型里的某些操作;版本号太高,目标推理引擎可能还没跟上。opset 11 是一个兼容性非常好的选择,对ResNet18这类经典模型支持完美。如果你用了特别新的网络结构,可能需要查阅文档,选择更高的版本。dynamic_axes:这个参数让模型支持动态形状。在上面的代码里,我把输入和输出的第0维(即batch size)标记为动态的,并命名为'batch_size'。这意味着转换出的ONNX模型,不仅可以处理batch size=1的输入,也可以处理batch size=2, 4, 8等。这在部署时非常有用,因为你的服务可能需要同时处理不同数量的请求。如果你确定batch size固定,可以省略这个参数。
转换完成后,我强烈建议你用ONNX Runtime验证一下转换是否正确。写个简单的脚本,分别用原始PyTorch模型和新的ONNX模型推理同一个输入,对比输出结果是否一致(允许微小的数值误差)。这是确保转换无误的最佳实践。
4. 第三步:深入提取模型权重为.wts文件
好了,现在我们有了ONNX这个通用格式。但对于某些特定的部署场景,比如使用NVIDIA的TensorRT进行极致优化时,我们可能需要一种更“原始”、更轻量的权重格式。.wts文件就是一种常见的纯文本权重存储格式,它简单到只记录每一层权重的名字和具体的数值。TensorRT的某些示例代码就使用这种格式来加载权重。
那么,如何从我们已经有的resnet18.pth文件中提取出这些权重,并打包成.wts文件呢?这个过程本质上就是遍历模型的状态字典(state_dict),把每个张量(Tensor)展平、转换成浮点数,然后按特定格式写入文本文件。
下面这个脚本是我参考了TensorRT官方样例后调整的,更加清晰和健壮,保存为export_to_wts.py。
import torch
import struct
def main():
# 1. 加载模型
model = torch.load('resnet18.pth', map_location='cpu')
model.eval()
# 2. 获取模型的状态字典(包含所有权重和偏置)
state_dict = model.state_dict()
# 3. 准备写入.wts文件
wts_filename = 'resnet18.wts'
with open(wts_filename, 'w') as f:
# 第一行写入权重的总条目数(即state_dict的key数量)
f.write(f"{len(state_dict)}\n")
# 4. 遍历状态字典中的每一项
for key, tensor in state_dict.items():
# 打印当前处理的层名称和形状,便于调试
print(f'Processing: {key}, shape: {tensor.shape}')
# 将权重张量展平为一维数组,并转换为CPU上的numpy数组(float32)
weight_data = tensor.cpu().numpy().flatten().astype('float32')
# 写入格式:<层名称> <权重数据长度>
f.write(f"{key} {len(weight_data)}")
# 5. 将每个浮点数转换为十六进制字符串并写入
# 使用 `struct.pack('>f', value)` 将float32打包为二进制,再用.hex()转为十六进制字符串
# '>' 表示大端字节序,这是一种常见的网络字节序,确保跨平台一致性
for value in weight_data:
hex_value = struct.pack('>f', value).hex()
f.write(f" {hex_value}")
# 每个层的数据写完后换行
f.write("\n")
print(f"\nAll weights have been exported to '{wts_filename}'")
print(f"Total layers processed: {len(state_dict)}")
if __name__ == '__main__':
main()
运行这个脚本,你会得到一个resnet18.wts的文本文件。用文本编辑器打开,它的结构是这样的:
第一行:总层数
第二行开始:层名1 权重数量1 十六进制数1 十六进制数2 ...
层名2 权重数量2 十六进制数1 十六进制数2 ...
...
这个过程有几个技术细节值得深究:
state_dict():这是PyTorch模型权重的“宝库”。它不是一个有序列表,而是一个字典(Dictionary),键(key)是每一层可学习参数的名字(如conv1.weight,bn1.bias),值(value)就是对应的权重张量。遍历它就能拿到所有权重。- 展平与类型转换:卷积层、全连接层的权重通常是多维的(如
[64, 3, 7, 7])。.flatten()操作把它们全部拉成一维长数组,方便顺序存储。.astype('float32')确保数据是单精度浮点数,这是最通用的格式。 - 十六进制存储:为什么用十六进制而不用直接的十进制小数?主要是为了精确和无歧义。浮点数在内存中以二进制形式存在,直接写小数到文本会有精度损失和字符串解析的问题。将其内存表示直接转为十六进制字符串,可以保证在另一端(如C++程序)能原封不动地还原出完全相同的二进制位,精度零损失。
- 字节序:
struct.pack('>f', value)中的>代表“大端序”(Big-endian)。这是一种约定俗成的字节排列顺序,在跨平台、跨语言的数据交换中,明确字节序可以避免很多诡异的错误。
导出了.wts文件,你就可以在TensorRT等引擎的C++代码中,写一个对应的解析器,把这些十六进制数读出来,还原成浮点数组,然后填充到构建好的网络层中。这就完成了从PyTorch训练到C++高性能推理的最后一环。
5. 转换过程中常见的“坑”与解决方案
走通了流程,不代表每次都能一帆风顺。模型转换这条路,我踩过的坑比走过的桥还多。下面我把几个最常见的问题和解决办法列出来,你遇到了可以直接来查。
5.1 ONNX转换失败:算子不支持
问题描述:运行torch.onnx.export时,报错提示某个算子(operator)不被当前opset版本支持,或者ONNX根本没有定义这个算子。
原因分析:PyTorch的算子集合非常庞大且活跃,ONNX的标准化过程有时会跟不上。特别是当你使用了一些比较新潮的、或者PyTorch特有的操作时。
解决方案:
- 降低opset版本:尝试将
opset_version从较高的版本(如13、14)降低到11或10。老版本更稳定,支持的算子虽然少,但都是久经考验的。 - 自定义算子:如果必须用某个算子,而ONNX又不支持,可以考虑自己实现一个“子图”来替代它。比如,用几个基础算子的组合来实现复杂算子的功能。这需要你对算子计算逻辑有深入理解。
- 修改模型结构:这是最根本但有时最有效的方法。在训练阶段,就尽量避免使用那些“稀奇古怪”的、可能不被广泛支持的层或操作。对于部署友好的模型设计,本身就是一个重要的课题。
- 查阅官方矩阵:PyTorch和ONNX官网通常有算子支持对照表,转换前可以先查一下。
5.2 动态形状导出问题
问题描述:转换ONNX时设置了动态维度(如动态batch size),但在推理引擎(如TensorRT)中加载时,却无法接受不同形状的输入。
原因分析:dynamic_axes参数设置不正确,或者推理引擎对动态形状的支持有特定要求(比如,可能只支持某些维度动态,或者需要额外的优化步骤)。
解决方案:
- 仔细检查
dynamic_axes:确保你指定的维度索引是正确的。例如,对于形状[batch, channel, height, width],第0维是batch。字典的键'input'必须和export函数中input_names里定义的名称完全一致。 - 分阶段转换:如果目标引擎对动态支持不好,可以考虑导出多个固定形状的ONNX模型,比如专门导出一个batch size=1的模型用于嵌入式设备,再导出一个batch size=8的模型用于服务器。
- 使用推理引擎的显式Batch模式:以TensorRT为例,它有一个“显式Batch”(Explicit Batch)维度的工作模式,在这种模式下定义网络,对动态形状的支持更好。这通常需要在构建TensorRT引擎时进行特定配置。
5.3 .wts文件解析错误
问题描述:在C++端读取.wts文件时,解析出来的权重值全是乱码或者NaN,导致模型推理结果完全错误。
原因分析:这几乎百分之百是数据对齐或解析逻辑的问题。要么是Python端写入的格式和C++端读取的格式不匹配,要么是字节序处理错了。
解决方案:
- 严格统一格式:确保C++解析代码的每一行逻辑,都和Python生成代码的写入逻辑镜像对应。比如,第一行读总数,然后循环读“名称、长度、数据”。一个空格或换行符的错位都会导致后续全部错乱。
- 验证字节序:确认C++端从十六进制字符串还原浮点数时,使用的字节序(大端
>还是小端<)与Python端(我们用了>)一致。在x86架构的PC上,本地字节序是小端,但我们的文件是大端序,所以解析时可能需要做转换,或者直接使用大端序解析函数。 - 写一个简单的验证脚本:在导出.wts文件后,可以马上写一个Python脚本重新读取它,将读出的权重与原始的
state_dict中的权重进行逐元素比较(允许极小的误差),确保写入和读取的过程是闭环正确的。把这个验证脚本作为流程的一部分,能提前发现很多问题。
5.4 精度下降与结果不一致
问题描述:PyTorch模型推理结果、ONNX模型推理结果、以及最终部署引擎上的推理结果,三者之间存在肉眼可见的数值差异,甚至导致分类错误。
原因分析:这是模型部署中最令人头疼的问题之一。原因可能多种多样:不同框架间算子实现的细微差异;浮点数计算顺序不同导致的累积误差;在转换或优化过程中某些层被融合或简化;甚至是在.wts文件中精度损失。
解决方案:
- 逐层对齐(Layer-wise Debugging):这是最有效的定位方法。不要只比较最终输出。尝试在PyTorch模型中插入钩子(hook),记录每一层(尤其是你怀疑的层,如BatchNorm、Softmax)的输出。在ONNX模型或最终部署模型中,也想办法获取对应层的输出。从输入开始,一层一层地对比,差异出现在哪一层,问题就大概率出在哪一层的转换或实现上。
- 关闭优化:一些推理引擎在导入ONNX时会默认进行图优化(如常量折叠、算子融合)。尝试关闭这些优化选项,让网络保持和原始ONNX一样的结构,先保证结果一致,再逐步开启优化排查。
- 检查数据预处理:确保在PyTorch、ONNX测试和最终部署中,输入数据的预处理流程完全一致。包括图像的归一化均值、标准差、缩放方式(是裁剪还是拉伸)、颜色通道顺序(RGB还是BGR)等。一个像素值的偏差,经过深度网络放大后,结果可能天差地别。
- 容忍合理误差:对于深度神经网络,由于浮点数计算的非结合性等原因,在不同平台、不同实现间产生
1e-5到1e-3量级的相对误差是正常的。只要最终分类的Top-1类别不变,通常可以接受。你需要确定一个可接受的误差阈值。
模型转换和部署是一个充满细节的工程活,它要求你对训练框架、中间格式、目标引擎都有一定的了解。但只要你耐心地、一步一步地走通这个流程,并学会了我上面提到的这些调试方法,你会发现,把训练好的模型成功地“搬”到任何需要它的地方,是一件非常有成就感的事情。
更多推荐


所有评论(0)