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

PyTorch实现无循环的多批次张量切片矩阵乘法

批量矩阵乘法实现(PyTorch)

现有两个PyTorch张量,形状定义如下:

A = [N x L x T]
B = [N x T x K]

需求:对每个N维度对应的切片执行矩阵乘法(如A[0,:,:] @ B[0,:,:],得到形状[L x K]的结果),最终输出形状为[N, L, K]的张量,且禁止通过循环遍历N维度以保证计算效率。


解决方案

方法1:直接使用torch.matmul

PyTorch的torch.matmul原生支持批量矩阵乘法,会自动识别并保留前导的批量维度(此处为N),对每个批量内的子张量执行矩阵乘法:

result = torch.matmul(A, B)
# 输出形状为[N, L, K],完全符合需求

方法2:使用torch.einsum(更直观的维度映射)

通过爱因斯坦求和符号明确指定维度对应关系,消除中间的T维度,直接得到目标形状:

result = torch.einsum('nlt,ntk->nlk', A, B)
  • 符号说明:nlt对应A的(N,L,T)维度,ntk对应B的(N,T,K)维度,箭头后的nlk表示输出维度为(N,L,K),其中T维度被求和抵消,与矩阵乘法逻辑一致。

验证示例

用小尺寸张量验证两种方法与循环结果的一致性:

import torch

N, L, T, K = 2, 3, 4, 5
A = torch.randn(N, L, T)
B = torch.randn(N, T, K)

# 循环实现(用于对比)
loop_result = torch.zeros(N, L, K)
for i in range(N):
    loop_result[i] = A[i] @ B[i]

# 方法1结果
matmul_result = torch.matmul(A, B)
# 方法2结果
einsum_result = torch.einsum('nlt,ntk->nlk', A, B)

# 检查一致性
print(torch.allclose(loop_result, matmul_result))  # 输出 True
print(torch.allclose(loop_result, einsum_result)) # 输出 True

内容的提问来源于stack exchange,提问作者Inkyu Kim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 05:15:11