在模型部署或二次开发过程中,我们有时需要修改模型权重文件(.pt 格式)中存储的标签名称(比如将英文标签改为中文、统一标签命名规范等)。本文将通过一段简单的代码,介绍如何使用 PyTorch 实现这一需求。
代码功能概述
这段代码的核心功能是:加载 PyTorch 模型权重文件(.pt),修改模型中存储的标签名称列表,最后将修改后的权重保存为新文件。适用于需要调整模型输出标签名称的场景(例如目标检测模型的类别标签修改)。
完整代码
import torch
# 加载模型
w = torch.load('D:/行人车辆.pt')
# 打印所有name
print(w.get('model').names)
# 定义一个将英文单词映射到中文单词的字典
word_map = {
'label': 'new_label',
}
# 遍历列表,将每个英文单词替换为其中文对应词
for i in range(len(w.get('model').names)):
if w.get('model').names[i] in word_map:
w.get('model').names[i] = word_map[w.get('model').names[i]]
# 打印替换后的列表
print('替换后')
print(w.get('model').names)
# 保存替换后的模型
torch.save(w, 'D:/new_best.pt')
代码分步解析
1. 导入必要库
首先需要导入 PyTorch 库,因为我们要处理.pt 格式的模型文件:
import torch
2. 加载模型权重文件
使用torch.load()方法加载本地的.pt 文件,这里需要注意文件路径的正确性(示例中为D:/行人车辆.pt):
w = torch.load('D:/行人车辆.pt')
加载后得到的w是一个字典对象,包含模型的权重、配置等信息。
3. 查看原始标签名称
模型的标签名称通常存储在权重字典的model.names字段中(不同模型可能有差异,需根据实际结构调整)。我们先打印原始标签,确认需要修改的内容:
print(w.get('model').names)
4. 定义标签映射关系
创建一个字典,键为原始标签名称,值为要替换的新标签名称。根据实际需求添加需要修改的标签对:
word_map = {
'label': 'new_label', # 示例:将'label'替换为'new_label'
# 可添加更多映射,如:'car': '车辆', 'person': '行人'
}
5. 遍历并替换标签
循环遍历标签列表,检查每个标签是否在映射字典中,若存在则替换为新标签:
for i in range(len(w.get('model').names)):
if w.get('model').names[i] in word_map:
w.get('model').names[i] = word_map[w.get('model').names[i]]
6. 验证修改结果
替换完成后,打印新的标签列表,确认修改是否符合预期:
print('替换后')
print(w.get('model').names)
7. 保存修改后的模型
使用torch.save()将修改后的权重保存为新文件(示例中为D:/new_best.pt),避免覆盖原始文件:
torch.save(w, 'D:/new_best.pt')
注意事项
- 文件路径:确保加载和保存的文件路径正确,避免因路径错误导致的文件无法读取 / 保存。
- 模型结构差异:不同模型的权重字典结构可能不同,
model.names并非通用字段,需根据实际模型的键名调整(可通过print(w.keys())查看字典结构)。 - 备份原始文件:修改前建议备份原始.pt 文件,防止操作失误导致文件损坏。
- 扩展映射关系:
word_map字典可根据实际需求扩展,支持批量替换多个标签。
4626

被折叠的 条评论
为什么被折叠?



