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
适用范围说明
两种方法都支持任意合法的批量矩阵乘法输入:
- 基础场景要求
a的形状为[B, M, K],b的形状为[B, N, K],即批量维度B相同,最后一维特征维度K相同 - 如果是多批量维度场景(例如
a形状为[B1, B2, M, K]),只需要先计算所有批量维度的总大小B = B1 * B2 * ...,其余逻辑不变即可适配 - 支持CPU、GPU等任意设备的张量,不需要移动数据
内容的提问来源于stack exchange,提问作者namespace-Pt
相关产品推荐
相关产品推荐

