PyTorch中批量实现卷积层输出与MLP标量权重相乘
批量场景下卷积输出与MLP权重的线性组合实现
核心是利用PyTorch的广播机制,让批量维度的权重能和每个卷积输出对应相乘,具体实现步骤如下:
调整MLP输出权重的形状
当前conv_weights形状为[128, 8],我们需要给它添加三个额外维度,用unsqueeze方法将其转为[128, 8, 1, 1, 1],这样就能和卷积输出的形状匹配,满足广播相乘的要求。堆叠卷积输出列表
把conv_outputs列表里的8个卷积输出(每个形状为[128, 32, 32, 32])用torch.stack在第1维度堆叠,得到形状为[128, 8, 32, 32, 32]的张量,这样每个样本对应的8个卷积输出就被整合到同一个批量条目下。加权求和得到最终结果
将调整后的权重与堆叠后的卷积输出相乘,然后在第1维度(卷积层维度)上求和,最终得到每个样本的线性组合结果,形状为[128, 32, 32, 32]。
完整的forward()实现代码:
def forward(self, x): # x.size() = [128, 3, 32, 32] conv_weights = self.mlp(x) # 输出形状 [128, 8] # 调整权重形状以适配广播 conv_weights = conv_weights.unsqueeze(1).unsqueeze(2).unsqueeze(3) # 变为 [128, 8, 1, 1, 1] # 执行所有卷积并收集输出 conv_outputs = [conv(x) for conv in self.convs] # 每个元素形状 [128, 32, 32, 32] # 堆叠卷积输出 stacked_convs = torch.stack(conv_outputs, dim=1) # 堆叠后形状 [128, 8, 32, 32, 32] # 加权相乘后求和得到最终结果 result = (stacked_convs * conv_weights).sum(dim=1) # 最终形状 [128, 32, 32, 32] return result
补充说明
torch.stack的dim=1参数是关键,它确保批量维度保持在第一位,后续的广播相乘能精准对应每个样本的权重和卷积输出。- 广播机制会自动将权重的单值扩展到卷积输出的通道、高度、宽度维度,无需手动复制数据,既简洁又高效。
内容的提问来源于stack exchange,提问作者James Wong
相关产品推荐
相关产品推荐

