PyTorch加载TinyYolo权重遇torch.load报错,Python3.10求解决方案
解决方案
你遇到的问题本质是官网下载的TinyYOLO权重是Darknet格式,并非PyTorch的state_dict文件,直接用torch.load()加载必然报错,和Python版本无关。以下是具体解决步骤:
1. 编写Darknet权重转PyTorch参数的函数
Darknet权重文件是二进制格式,需要按模型层的顺序逐个读取参数并赋值给PyTorch模型的对应层。添加以下转换函数:
import numpy as np import torch import torch.nn as nn def load_darknet_weights(model, weights_path): # 打开权重文件 with open(weights_path, 'rb') as f: # 跳过前5个int32参数(Darknet的头信息) np.fromfile(f, dtype=np.int32, count=5) # 遍历模型的卷积层和BN层,加载权重 for module in model.cnn: if isinstance(module, nn.Conv2d): # 加载卷积层权重 num_weights = module.weight.numel() conv_weights = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_weights)) conv_weights = conv_weights.view_as(module.weight) module.weight.data.copy_(conv_weights) # 如果卷积层有偏置(你的conv8设置了bias=False,跳过) if module.bias is not None: num_bias = module.bias.numel() bias_weights = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_bias)) bias_weights = bias_weights.view_as(module.bias) module.bias.data.copy_(bias_weights) elif isinstance(module, nn.BatchNorm2d): # 加载BN层的bias、weight、running_mean、running_var num_bias = module.bias.numel() bias = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_bias)) bias = bias.view_as(module.bias) module.bias.data.copy_(bias) num_weight = module.weight.numel() weight = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_weight)) weight = weight.view_as(module.weight) module.weight.data.copy_(weight) num_running_mean = module.running_mean.numel() running_mean = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_running_mean)) running_mean = running_mean.view_as(module.running_mean) module.running_mean.data.copy_(running_mean) num_running_var = module.running_var.numel() running_var = torch.from_numpy(np.fromfile(f, dtype=np.float32, count=num_running_var)) running_var = running_var.view_as(module.running_var) module.running_var.data.copy_(running_var) return model
2. 修改load_model函数
替换原来直接用torch.load()的逻辑,改用上面的转换函数:
def load_model(weights): model = TinyYoloNet() # 调用转换函数加载Darknet权重 model = load_darknet_weights(model, weights) return model.cuda() if use_gpu else model # 加载权重 model = load_model(weights='weights/yolov2-tiny-voc.weights')
3. 关键注意事项
- 必须确保你的
TinyYoloNet类中所有卷积层、BN层的顺序、输出通道数、卷积核大小,完全和Darknet官方TinyYOLOv2的结构一致(你代码中注释掉的conv1到conv7需要补全正确结构,否则权重对应会出错) - 若加载后仍出现参数不匹配的报错,检查模型每层的配置是否和Darknet官方定义对齐
内容的提问来源于stack exchange,提问作者khatooon khedri
相关产品推荐
相关产品推荐

