You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中加载.mat格式的NetVLAD预训练模型?

解决方法:将MatLab训练的NetVLAD模型迁移到PyTorch

不能直接用load_state_dict加载.mat格式的模型,因为MatLab导出的参数结构、维度顺序和PyTorch的state_dict完全不匹配,必须手动进行参数映射和维度转换。以下是具体步骤:

1. 先分析.mat文件的参数结构

首先打印并查看.mat文件里的参数组织方式,确认每个层的权重、偏置以及NetVLAD相关组件(聚类中心、白化矩阵)的存储位置:

from scipy.io import loadmat
import numpy as np

data = loadmat('vd16_pitts30k_conv5_3_vlad_preL2_intra_white.mat')
net_data = data['net']

# 打印参数结构,根据输出调整后续的参数提取逻辑
print(net_data.keys())
# 若为struct类型,可遍历查看内部字段,比如:
# print(net_data['params'].shape)
# print(net_data['params'][0])

2. 定义与原模型匹配的PyTorch版NetVLAD结构

需要复现原MatLab模型的网络结构(VGG16的conv5_3层 + NetVLAD模块 + 白化层):

import torch
import torch.nn as nn
import torch.nn.functional as F


class NetVLAD(nn.Module):
    def __init__(self, num_clusters=64, dim=512, normalize_input=True):
        super().__init__()
        self.num_clusters = num_clusters
        self.dim = dim
        self.normalize_input = normalize_input
        self.conv = nn.Conv2d(dim, num_clusters, kernel_size=1, bias=True)
        self.centroids = nn.Parameter(torch.rand(num_clusters, dim))

    def forward(self, x):
        N, C, H, W = x.shape
        if self.normalize_input:
            x = F.normalize(x, p=2, dim=1)
        
        soft_assign = self.conv(x).view(N, self.num_clusters, -1)
        soft_assign = F.softmax(soft_assign, dim=1)
        
        x_flatten = x.view(N, C, -1)
        vlad = torch.zeros([N, self.num_clusters, C], dtype=x.dtype, device=x.device)
        for Ck in range(self.num_clusters):
            residual = x_flatten - self.centroids[Ck:Ck+1, :].unsqueeze(2)
            residual *= soft_assign[:, Ck:Ck+1, :]
            vlad[:, Ck:Ck+1, :] = residual.sum(dim=2)
        
        vlad = F.normalize(vlad, p=2, dim=2)
        vlad = vlad.view(N, -1)
        vlad = F.normalize(vlad, p=2, dim=1)
        return vlad


class VGG16NetVLAD(nn.Module):
    def __init__(self, num_clusters=64):
        super().__init__()
        # 构建VGG16到conv5_3的特征提取层
        vgg = torch.hub.load('pytorch/vision:v0.10.0', 'vgg16', pretrained=False)
        self.features = nn.Sequential(*list(vgg.features.children())[:-3])
        self.netvlad = NetVLAD(num_clusters=num_clusters, dim=512)
        # 原模型包含白化层,需添加对应结构
        self.whitening = nn.Linear(64*512, 64*512, bias=False)

    def forward(self, x):
        x = self.features(x)
        x = self.netvlad(x)
        x = self.whitening(x)
        x = F.normalize(x, p=2, dim=1)
        return x

3. 手动映射MatLab参数到PyTorch模型

根据第一步分析的参数结构,将.mat中的参数转换维度后赋值给PyTorch模型的对应层:

# 初始化模型
model = VGG16NetVLAD(num_clusters=64)

# 加载VGG特征层的卷积参数
# 假设.mat中VGG的conv层权重按顺序存储在net_data['params']的前若干位置
for idx, layer in enumerate(model.features):
    if isinstance(layer, nn.Conv2d):
        # MatLab卷积权重维度:[H, W, 输入通道数, 输出通道数]
        # PyTorch卷积权重维度:[输出通道数, 输入通道数, H, W]
        mat_weight = net_data['params'][idx*2]
        mat_bias = net_data['params'][idx*2 + 1]
        
        pt_weight = torch.from_numpy(np.transpose(mat_weight, (3, 2, 0, 1)))
        pt_bias = torch.from_numpy(mat_bias.squeeze())
        
        layer.weight.data.copy_(pt_weight)
        layer.bias.data.copy_(pt_bias)

# 加载NetVLAD模块的参数
# 分配层的卷积参数
mat_vlad_conv_weight = net_data['params'][-4]
mat_vlad_conv_bias = net_data['params'][-3]
model.netvlad.conv.weight.data.copy_(torch.from_numpy(np.transpose(mat_vlad_conv_weight, (3, 2, 0, 1))))
model.netvlad.conv.bias.data.copy_(torch.from_numpy(mat_vlad_conv_bias.squeeze()))

# 聚类中心参数
mat_centroids = net_data['params'][-2]
model.netvlad.centroids.data.copy_(torch.from_numpy(np.transpose(mat_centroids, (1, 0))))

# 加载白化层参数
mat_whitening = net_data['params'][-1]
model.whitening.weight.data.copy_(torch.from_numpy(np.transpose(mat_whitening, (1, 0))))

关键注意点

  • 维度转换:MatLab与PyTorch的卷积权重维度顺序差异极大,必须通过np.transpose调整后才能使用。
  • 参数对应:不同版本的.mat文件参数存储顺序可能有差异,必须通过第一步的结构分析确认每个参数对应的模型层。
  • 白化层:原模型包含intra_white白化操作,需在PyTorch模型中添加对应的线性层并加载参数。

内容的提问来源于stack exchange,提问作者Shania F.

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 03:38:11