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

如何用Weights & Biases管理PyTorch CNN架构实验并实现结果复现?

问题:如何用W&B完整复现包含模型类代码的CNN实验?

我跟着PyTorch官方基础教程做实验,尝试不同的CNN架构(调整层数、每层通道数等),想用W&B规范管理实验流程。目前能保存模型权重和参数,但没法保存模型类的代码,导致实验结果很难复现。我考虑过用inspect模块保存源码,但不确定这是不是最优方案,而且Google Colab不支持这个模块。附上现有代码,想问问有没有更优的实验管理方法,能实现包含模型类代码的结果复现?

现有代码

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

torch.manual_seed(0)

# MODEL DESCRIPTION
class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = torch.flatten(x, 1) # flatten all dimensions except batch
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

model = Net()
model_desc = repr(model)

# TRAINING....

# EVALUATING...
accuracy = ...

# SAVE RESULTS
run = wandb.init(project='test-project', name=model_desc)

artifact = wandb.Artifact('model', type='model')
artifact.add_file(MODEL_PATH)

# 2. Save mode inputs and hyperparameters
config = run.config
config.test_number = 1
config.model_desc = model_desc

# 3. Log metrics over time to visualize performance
run.log({"accuracy": accuracy})

# 4. Log an artifact to W&B
run.log_artifact(artifact)
run.finish()

可行解决方案

1. 手动上传模型代码文件到W&B Artifact

如果是在Colab中运行,可以先把模型类代码写入临时文件,再将该文件添加到Artifact中,确保源码随实验一起保存:

# 在Colab中写入模型代码到临时文件
with open("model_def.py", "w") as f:
    f.write("""
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = torch.flatten(x, 1)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x
""")

# 将模型代码文件加入Artifact
artifact.add_file("model_def.py")

2. 开启W&B自动代码保存功能

在初始化W&B run时设置save_code=True,W&B会自动保存当前运行的脚本(Colab环境下会保存Notebook的快照),直接关联到对应实验,无需手动处理:

run = wandb.init(project='test-project', name=model_desc, save_code=True)

之后在W&B的实验详情页,就能找到Files标签页查看保存的代码快照,完全复现当时的模型定义。

3. 参数化模型结构,用配置驱动模型构建

把模型的关键结构参数(如卷积通道数、全连接层大小、层数等)提取到W&B配置中,动态构建模型。这样只需保存配置参数,就能复现模型结构,不需要单独保存类代码:

# 初始化W&B时定义模型结构配置
run = wandb.init(project='test-project', name='parametrized-cnn', 
                 config={
                     "conv_channels": [3, 6, 16],
                     "conv_kernels": [5, 5],
                     "fc_sizes": [120, 84, 10],
                     "pool_size": 2
                 })
config = run.config

# 基于配置动态构建模型
class Net(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.conv_layers = nn.ModuleList()
        in_channels = config.conv_channels[0]
        for out_channels in config.conv_channels[1:]:
            self.conv_layers.append(nn.Conv2d(in_channels, out_channels, config.conv_kernels[0]))
            in_channels = out_channels
        self.pool = nn.MaxPool2d(config.pool_size, config.pool_size)
        
        # 计算全连接层输入维度(假设输入是32x32的CIFAR-10图片)
        self.fc_input_size = config.conv_channels[-1] * (32 // (2**len(config.conv_channels[1:])))**2
        self.fc_layers = nn.ModuleList()
        in_features = self.fc_input_size
        for out_features in config.fc_sizes[:-1]:
            self.fc_layers.append(nn.Linear(in_features, out_features))
            in_features = out_features
        self.fc_final = nn.Linear(in_features, config.fc_sizes[-1])

    def forward(self, x):
        for conv in self.conv_layers:
            x = self.pool(F.relu(conv(x)))
        x = torch.flatten(x, 1)
        for fc in self.fc_layers:
            x = F.relu(fc(x))
        x = self.fc_final(x)
        return x

model = Net(config)

这种方式下,W&B会自动保存配置参数,复现时只需加载配置就能重建完全一致的模型。


内容的提问来源于stack exchange,提问作者Ariel Yael

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 15:02:24