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

PyTorch如何从普通矩阵乘积结果中高效提取批量矩阵乘的结果

批量矩阵乘法结果提取方法

你观察到的现象本质是:当你把形状为[B, M, K]的批量张量a和形状为[B, N, K]的批量张量b的前两维展平后做全连接矩阵乘法,得到的g形状为[B*M, B*N],其中每个batch对应的有效批量乘法结果,是g中沿对角线分布的M*N大小的块,其余非对角块是跨batch的无效乘积结果。

以下是两种零拷贝、可泛化的高效提取方案,所有操作均返回张量视图,无额外内存拷贝,效率和原生切片一致:

方法1:reshape+对角提取(代码最简洁)

import torch

def extract_c_from_g(g, a_shape, b_shape):
    # 解析维度:a的形状为[B, M, K],b的形状为[B, N, K]
    B, M, K = a_shape
    B, N, K = b_shape
    # 先将g重构为[B, M, B, N]的四维张量
    g_reshaped = g.reshape(B, M, B, N)
    # 提取dim0和dim2的对角元素(即同一batch的块),调整维度顺序得到[B, M, N]
    c = g_reshaped.diagonal(dim1=0, dim2=2).permute(2, 0, 1)
    return c

示例验证

# 你的示例输入
a = torch.arange(12, dtype=torch.float).view(2,3,2)
b = torch.arange(12, dtype=torch.float).view(2,3,2) - 1
c_true = a.matmul(b.transpose(-1,-2))
e = a.view(6,2)
f = b.view(6,2)
g = e.matmul(f.transpose(-1,-2))

# 调用函数提取
c_extracted = extract_c_from_g(g, a.shape, b.shape)
# 验证结果一致
print(torch.allclose(c_extracted, c_true)) # 输出True

方法2:高级索引(逻辑更直观)

def extract_c_from_g_v2(g, a_shape, b_shape):
    B, M, K = a_shape
    B, N, K = b_shape
    # 生成和g同设备的行、列分块索引
    row_idx = torch.arange(B*M, device=g.device).view(B, M)
    col_idx = torch.arange(B*N, device=g.device).view(B, N)
    # 广播索引提取对应块,直接得到[B, M, N]形状的结果
    c = g[row_idx.unsqueeze(-1), col_idx.unsqueeze(1)]
    return c

适用范围说明

两种方法都支持任意合法的批量矩阵乘法输入:

  1. 基础场景要求a的形状为[B, M, K],b的形状为[B, N, K],即批量维度B相同,最后一维特征维度K相同
  2. 如果是多批量维度场景(例如a形状为[B1, B2, M, K]),只需要先计算所有批量维度的总大小B = B1 * B2 * ...,其余逻辑不变即可适配
  3. 支持CPU、GPU等任意设备的张量,不需要移动数据

内容的提问来源于stack exchange,提问作者namespace-Pt

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 20:24:03