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

PyTorch中如何高效对输入张量不同区域应用不同MLP

按索引选择MLP处理输入的高效方案

问题背景

我有一个形状为(8,10)的输入张量input_data,同时定义了三个结构完全一致、输入尺寸为10的MLP(mlp1、mlp2、mlp3)。另外有一个形状为(8,)的索引张量mlp_index,用于指定输入的每一行要使用哪个MLP进行计算(例如mlp_index[0]=2时,对input_data[0]应用mlp3)。

我尝试了几种处理方式,但发现仅用单个MLP处理整个输入的速度显著更快,希望找到更高效的多MLP选择计算方案。

测试代码及结果

示例代码

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

torch.manual_seed(42)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

class MLP0(nn.Module):
    def __init__(self, input_size, hidden_size, output_size):
        super(MLP0, self).__init__()
        self.fc1 = nn.Linear(input_size, hidden_size)
        self.fc2 = nn.Linear(hidden_size, output_size)

    def forward(self, x):
        x = F.relu(self.fc1(x))
        x = self.fc2(x)
        return x

# 模型参数
input_size = 10
hidden_size = 20
output_size = 5

mlp1 = MLP0(input_size, hidden_size, output_size).to(device)
mlp2 = MLP0(input_size, hidden_size, output_size).to(device)
mlp3 = MLP0(input_size, hidden_size, output_size).to(device)

input_data = torch.rand(size=(8, input_size), device=device)
mlp_index = torch.tensor([0, 1, 0, 1, 0, 2, 0, 2], device=device)

# 基准:单MLP处理全部输入
def baseline():
    return mlp1(input_data)

# 方法1:先计算所有MLP的输出再用where选择
def first_update():
    out_1 = mlp1(input_data)
    out_2 = mlp2(input_data)
    out_3 = mlp3(input_data)
    result = torch.where(mlp_index == 0, out_1, 
                         torch.where(mlp_index == 1, out_2, out_3))
    return result

# 方法2:在where中直接调用MLP
def second_update():
    result = torch.where(mlp_index == 0, mlp1(input_data), 
                         torch.where(mlp_index == 1, mlp2(input_data), mlp3(input_data)))
    return result

# 方法3:按索引拆分输入,分别计算后拼接
def third_update():
    mask1 = mlp_index == 0
    mask2 = mlp_index == 1
    mask3 = mlp_index == 2
    
    out_1 = mlp1(input_data[mask1])
    out_2 = mlp2(input_data[mask2])
    out_3 = mlp3(input_data[mask3])
    
    out = torch.zeros(size=(8, output_size), device=device)
    out[mask1] = out_1
    out[mask2] = out_2
    out[mask3] = out_3
    return out

# 测试耗时
baseline_time = timeit.timeit(baseline, number=20000)
print(f"基准单MLP耗时: {baseline_time:.4f} 秒")

first_update_time = timeit.timeit(first_update, number=20000)
print(f"方法1耗时: {first_update_time:.4f} 秒")

second_update_time = timeit.timeit(second_update, number=20000)
print(f"方法2耗时: {second_update_time:.4f} 秒")

third_update_time = timeit.timeit(third_update, number=20000)
print(f"方法3耗时: {third_update_time:.4f} 秒")

测试输出

基准单MLP耗时: 1.5391 秒
方法1耗时: 2.1762 秒
方法2耗时: 2.2332 秒
方法3耗时: 6.2527 秒

高效解决方案:合并MLP参数批量计算

由于三个MLP结构完全一致,只是参数不同,我们可以将它们的参数合并成批量张量,然后根据索引选择对应参数进行一次前向传播,避免多次调用MLP的开销。

实现代码

def merged_mlp_forward():
    # 合并三个MLP的fc1参数:权重形状变为(3, hidden_size, input_size),偏置变为(3, hidden_size)
    fc1_weights = torch.stack([mlp1.fc1.weight, mlp2.fc1.weight, mlp3.fc1.weight])
    fc1_biases = torch.stack([mlp1.fc1.bias, mlp2.fc1.bias, mlp3.fc1.bias])
    
    # 合并fc2参数:权重形状(3, output_size, hidden_size),偏置(3, output_size)
    fc2_weights = torch.stack([mlp1.fc2.weight, mlp2.fc2.weight, mlp3.fc2.weight])
    fc2_biases = torch.stack([mlp1.fc2.bias, mlp2.fc2.bias, mlp3.fc2.bias])
    
    # 选择每个样本对应的参数
    selected_fc1_w = fc1_weights[mlp_index]  # 形状(8, hidden_size, input_size)
    selected_fc1_b = fc1_biases[mlp_index]  # 形状(8, hidden_size)
    selected_fc2_w = fc2_weights[mlp_index]  # 形状(8, output_size, hidden_size)
    selected_fc2_b = fc2_biases[mlp_index]  # 形状(8, output_size)
    
    # 批量计算:先做fc1的线性变换 + ReLU
    hidden = torch.bmm(selected_fc1_w, input_data.unsqueeze(-1)).squeeze(-1) + selected_fc1_b
    hidden = F.relu(hidden)
    
    # 再做fc2的线性变换
    output = torch.bmm(selected_fc2_w, hidden.unsqueeze(-1)).squeeze(-1) + selected_fc2_b
    return output

# 测试耗时
merged_time = timeit.timeit(merged_mlp_forward, number=20000)
print(f"合并参数批量计算耗时: {merged_time:.4f} 秒")

效果说明

这种方法把多次MLP调用转换成一次批量张量运算,避免了模型调用的额外开销,同时也省去了数据拆分、拼接的操作。实际测试中,它的耗时会非常接近单MLP的基准耗时,大幅优于之前的三种方法。

补充优化

如果需要频繁进行这类操作,可以将合并参数的逻辑封装成一个自定义模块,避免每次前向都重复堆叠参数:

class MergedMLP(nn.Module):
    def __init__(self, mlps):
        super().__init__()
        # 提取并堆叠所有MLP的参数
        self.fc1_weights = nn.Parameter(torch.stack([m.fc1.weight for m in mlps]))
        self.fc1_biases = nn.Parameter(torch.stack([m.fc1.bias for m in mlps]))
        self.fc2_weights = nn.Parameter(torch.stack([m.fc2.weight for m in mlps]))
        self.fc2_biases = nn.Parameter(torch.stack([m.fc2.bias for m in mlps]))
        
    def forward(self, x, indices):
        selected_fc1_w = self.fc1_weights[indices]
        selected_fc1_b = self.fc1_biases[indices]
        selected_fc2_w = self.fc2_weights[indices]
        selected_fc2_b = self.fc2_biases[indices]
        
        hidden = torch.bmm(selected_fc1_w, x.unsqueeze(-1)).squeeze(-1) + selected_fc1_b
        hidden = F.relu(hidden)
        output = torch.bmm(selected_fc2_w, hidden.unsqueeze(-1)).squeeze(-1) + selected_fc2_b
        return output

# 使用示例
merged_mlp = MergedMLP([mlp1, mlp2, mlp3]).to(device)

def merged_module_forward():
    return merged_mlp(input_data, mlp_index)

merged_module_time = timeit.timeit(merged_module_forward, number=20000)
print(f"自定义合并模块耗时: {merged_module_time:.4f} 秒")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 12:24:51