求助:VGG16权重numpy object数组整体校验失败但逐个校验通过
VGG16权重校验异常问题解答
核心原因
问题就出在numpy对dtype=object数组的对比逻辑上。
当你用dtype=object创建numpy数组时,数组里存的是各个子numpy数组的内存引用,而非子数组的实际数值。np.array_equal(params1, params2)会直接对比这些引用是否指向同一个内存对象——两次加载模型生成的子数组是不同的内存实例,所以返回False。
而你逐个迭代对比时,np.array_equal(val1, val2)是直接校验子数组的数值内容,两次加载的参数数值完全一致,所以结果全为True。
验证引用差异
用下面的代码就能确认这一点:
print(params1[0] is params2[0]) # 输出False,说明是不同内存对象 print(np.array_equal(params1[0], params2[0])) # 输出True,数值内容一致
正确的校验方法
如果你要整体校验参数内容,别直接用np.array_equal,可以选这两种方式:
- 逐个元素校验后确认全部一致:
all_equal = all(np.array_equal(p1, p2) for p1, p2 in zip(params1, params2)) print(all_equal) # 输出True - 改用Python列表存储参数数组,列表的
==会递归对比元素内容:params1_list = [param.detach().numpy() for param in vgg.parameters()] params2_list = [param.detach().numpy() for param in vgg2.parameters()] print(params1_list == params2_list) # 输出True
更高效的PyTorch原生校验方式
其实没必要转numpy数组,PyTorch本身就有直接的参数对比方法:
all_params_equal = all(torch.equal(p1, p2) for p1, p2 in zip(vgg.parameters(), vgg2.parameters())) print(all_params_equal) # 输出True
这种方式既高效,又能避免numpy object数组带来的坑。
内容的提问来源于stack exchange,提问作者Nikaido
相关产品推荐
相关产品推荐

