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
相关产品推荐
相关产品推荐

