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

PyTorch中批量实现卷积层输出与MLP标量权重相乘

批量场景下卷积输出与MLP权重的线性组合实现

核心是利用PyTorch的广播机制,让批量维度的权重能和每个卷积输出对应相乘,具体实现步骤如下:

  1. 调整MLP输出权重的形状
    当前conv_weights形状为[128, 8],我们需要给它添加三个额外维度,用unsqueeze方法将其转为[128, 8, 1, 1, 1],这样就能和卷积输出的形状匹配,满足广播相乘的要求。

  2. 堆叠卷积输出列表
    把conv_outputs列表里的8个卷积输出(每个形状为[128, 32, 32, 32])用torch.stack在第1维度堆叠,得到形状为[128, 8, 32, 32, 32]的张量,这样每个样本对应的8个卷积输出就被整合到同一个批量条目下。

  3. 加权求和得到最终结果
    将调整后的权重与堆叠后的卷积输出相乘,然后在第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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 15:55:20