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

如何高效训练多个初始化不同的同架构小型神经网络?

高效训练多个同架构不同初始化的小模型

我需要在同一训练数据上训练多个仅初始化不同、架构完全相同的极小模型,内存完全能同时容纳这些模型。但朴素写法的训练时间随模型数量线性增长,虽然CUDA支持非阻塞式计算和自动并行化,但没发挥出优势。

朴素实现代码

import time
import numpy as np
import torch
import torch.nn as nn

class MLP(nn.Module):
    def __init__(self, network_size):
        super(MLP, self).__init__()
        self.fc1 = nn.Linear(2, network_size)
        self.fc2 = nn.Linear(network_size, 1)

    def forward(self, x):
        return torch.sigmoid(self.fc2(self.fc1(x)))

def train(num_networks, network_size, num_iterations):
    criterion = torch.nn.BCELoss()
    data = torch.zeros((5, 2), device='cuda')
    targets = torch.ones((5, 1), device='cuda')
    
    models = []
    for _ in range(num_networks):
        models.append(MLP(network_size).cuda())
    for model in models:
        optimizer = torch.optim.Adam(model.parameters())
        for _ in range(num_iterations):
            output = model(data)
            loss = criterion(output, targets)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

training_start = time.perf_counter()
train(1, 20, 1000)
print(f"Training 1 model took {time.perf_counter() - training_start:.2f}s")

training_start = time.perf_counter()
train(5, 20, 1000)
print(f"Training 5 models took {time.perf_counter() - training_start:.2f}s")

training_start = time.perf_counter()
train(15, 20, 1000)
print(f"Training 15 models took {time.perf_counter() - training_start:.2f}s")

朴素实现输出

Training 1 model took 0.68s
Training 5 models took 3.36s
Training 15 models took 10.18s

合并模型优化实现

我通过将多个模型合并为一个大网络实现了效率提升,但这种方式存在易出错的问题——比如调整网络规模、提取单个训练好的模型时都很麻烦。

合并模型代码

class MergedMLP(nn.Module):
    def __init__(self, num_networks, network_size):
        super().__init__()
        self.fc1 = nn.Linear(2, num_networks * network_size, device='cuda')
        self.fc2 = nn.Linear(num_networks * network_size, num_networks, device='cuda')
        
        self.fc2_weight_mask = torch.zeros_like(self.fc2.weight.data, device='cuda', requires_grad=False)
        for i in range(num_networks):
            self.fc2_weight_mask[i,i*network_size:(i+1)*network_size] = 1
        self.fc2.weight.data *= self.fc2_weight_mask

    def forward(self, x):
        return torch.sigmoid(self.fc2(self.fc1(x)))
    
def train_merged(num_networks, network_size, num_iterations):
    criterion = torch.nn.BCELoss()
    data = torch.zeros((5, 2), device='cuda')
    targets = torch.ones((5, num_networks), device='cuda')
    
    model = MergedMLP(num_networks, network_size).cuda()
    optimizer = torch.optim.Adam(model.parameters())
    for _ in range(num_iterations):
        output = model(data)
        loss = criterion(output, targets)
        optimizer.zero_grad()
        loss.backward()
        model.fc2.weight.grad *= model.fc2_weight_mask
        optimizer.step()
        
training_start = time.perf_counter()
train_merged(15, 20, 1000)
print(f"Training merged models took {time.perf_counter() - training_start:.2f}s")

合并模型输出

Training merged models took 0.70s

核心疑问

能不能用更接近朴素实现的代码达到和合并模型一样的运行效率?合并实现的可维护性太差,调整和拆分都很容易出错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 22:45:01