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

基于VAE的CNN模型参数降维与重构的PyTorch实现问询

实现CNN参数的VAE编码与重构流程

步骤1:提取并扁平化CNN参数

首先需要把训练好的CNN模型参数提取出来,拼成一维张量,同时记录每个参数的形状与名称,方便后续重构。

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

# 原CNN模型(用户提供)
class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = nn.Conv2d(1, 10, kernel_size=5)
        self.conv2 = nn.Conv2d(10, 20, kernel_size=5)
        self.conv2_drop = nn.Dropout2d()
        self.fc1 = nn.Linear(320, 50)
        self.fc2 = nn.Linear(50, 10)

    def forward(self, x):
        x = F.relu(F.max_pool2d(self.conv1(x), 2))
        x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 2))
        x = x.view(-1, 320)
        x = F.relu(self.fc1(x))
        x = F.dropout(x, training=self.training)
        x = self.fc2(x)
        return F.log_softmax(x, dim=1)  # 补充dim参数避免警告

# 实例化并加载训练好的CNN权重
cnn_model = Net()
# 模拟加载训练完成的权重(实际使用时替换为真实路径)
# cnn_model.load_state_dict(torch.load('trained_cnn_weights.pth'))

# 提取参数信息并扁平化
param_names = []
param_shapes = []
flat_param_list = []

for name, param in cnn_model.named_parameters():
    param_names.append(name)
    param_shapes.append(param.shape)
    flat_param_list.append(param.flatten())

# 拼接为带batch维度的一维张量
flat_param_tensor = torch.cat(flat_param_list).unsqueeze(0)
total_params = flat_param_tensor.size(1)
print(f"CNN总参数数量:{total_params}")

步骤2:修改VAE适配参数向量输入

原VAE是为图像设计的,需要改为处理一维参数向量的结构,编码器和解码器改用全连接层:

# 适配参数向量的VAE模型
class ParamVAE(nn.Module):
    def __init__(self, input_dim, h_dim=512, z_dim=32):
        super(ParamVAE, self).__init__()
        # 编码器:将一维参数向量编码到隐空间
        self.encoder = nn.Sequential(
            nn.Linear(input_dim, h_dim),
            nn.ReLU(),
            nn.Linear(h_dim, h_dim//2),
            nn.ReLU(),
        )
        # 均值与方差输出层
        self.fc_mu = nn.Linear(h_dim//2, z_dim)
        self.fc_logvar = nn.Linear(h_dim//2, z_dim)
        
        # 解码器:从隐空间重构参数向量
        self.decoder = nn.Sequential(
            nn.Linear(z_dim, h_dim//2),
            nn.ReLU(),
            nn.Linear(h_dim//2, h_dim),
            nn.ReLU(),
            nn.Linear(h_dim, input_dim),
            nn.Sigmoid()  # 若参数范围大,可先归一化再用此激活,否则可移除
        )
        
    def reparameterize(self, mu, logvar):
        std = torch.exp(0.5 * logvar)
        eps = torch.randn_like(std)
        return mu + eps * std
    
    def encode(self, x):
        h = self.encoder(x)
        mu = self.fc_mu(h)
        logvar = self.fc_logvar(h)
        z = self.reparameterize(mu, logvar)
        return z, mu, logvar
    
    def decode(self, z):
        return self.decoder(z)
    
    def forward(self, x):
        z, mu, logvar = self.encode(x)
        recon_x = self.decode(z)
        return recon_x, mu, logvar

步骤3:训练VAE(以CNN参数为训练数据)

用CNN参数作为训练数据,让VAE学习编码和解码的映射关系:

# 实例化VAE
vae = ParamVAE(input_dim=total_params)
optimizer = torch.optim.Adam(vae.parameters(), lr=1e-3)
recon_criterion = nn.MSELoss()  # 参数回归用MSE损失

# 训练循环(示例用单个CNN参数,实际建议收集多个同结构CNN参数作为数据集)
epochs = 1000
for epoch in range(epochs):
    optimizer.zero_grad()
    recon_params, mu, logvar = vae(flat_param_tensor)
    
    # 重构损失
    recon_loss = recon_criterion(recon_params, flat_param_tensor)
    # KL散度损失
    kl_loss = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
    # 总损失(可调整KL损失权重)
    total_loss = recon_loss + kl_loss * 0.01
    
    total_loss.backward()
    optimizer.step()
    
    if (epoch + 1) % 100 == 0:
        print(f"Epoch {epoch+1}, 总损失: {total_loss.item():.4f}, 重构损失: {recon_loss.item():.4f}")

步骤4:编码参数到隐空间并重构CNN模型

训练完成后,即可完成参数的编码-解码,并将重构参数加载回CNN模型:

# 编码与解码(禁用梯度计算)
with torch.no_grad():
    z, mu, logvar = vae.encode(flat_param_tensor)
    print(f"隐空间向量维度: {z.shape}")
    recon_flat_params = vae.decode(z)

# 将重构的扁平化参数恢复为原CNN参数形状
recon_cnn = Net()
start_idx = 0
for name, shape in zip(param_names, param_shapes):
    param_numel = shape.numel()
    # 提取对应片段并恢复形状
    recon_param = recon_flat_params[0, start_idx:start_idx+param_numel].view(shape)
    # 加载到重构模型
    recon_cnn.state_dict()[name].copy_(recon_param)
    start_idx += param_numel

# recon_cnn即为VAE重构参数得到的CNN模型

关键注意点

  • 参数归一化:若CNN参数数值范围较大,建议先将参数归一化到0-1或-1到1区间,训练后再反归一化,提升训练效果。
  • 数据集扩展:单一样本训练的VAE泛化能力差,建议收集多个同结构、不同训练状态的CNN参数(如不同epoch权重、不同初始化的训练模型)作为数据集。
  • 结构调整:可根据参数总维度调整VAE的隐藏层维度h_dim和隐空间维度z_dim,找到合适的降维比例。

内容的提问来源于stack exchange,提问作者Atefe Hassani

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 04:09:51