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

批次索引不相等的批量矩阵与向量乘积实现方案问询

高效实现跨批次的批量矩阵/矩阵-向量乘积(支持CPU/GPU)

问题背景

需要实现两种跨批次的张量运算,要求同时支持CPU和GPU,且避免冗余内存分配:

  1. 批量矩阵-矩阵乘积:给定形状为(m×n×N_a)的张量A和(n×p×N_b)的张量B,输出C[:,:,ia,ib] = A[:,:,ia] * B[:,:,ib](最终形状(m×p×N_a×N_b))
  2. 批量矩阵-向量乘积:给定形状为(m×n×N_a)的张量A和(n×N_b)的张量B,输出D[:,ia,ib] = A[:,:,ia] * B[:,ib](最终形状(m×N_a×N_b))

此前用Tullio.jl可实现需求,但该库已无法兼容新版CUDA.jl;尝试NNlib.jl的batched_mul/batched_vec时,需大量repeat操作对齐批次索引,带来冗余计算和内存开销。希望找到更高效的实现方式,且避免使用KernelAbstractions(需自行定义梯度)。

解决方案

以下两种方案均支持自动微分(无需手动定义梯度),且同时兼容CPU与GPU:

方案1:TensorOperations.jl(直观简洁)

TensorOperations.jl基于张量收缩的数学定义实现运算,天然支持CPU/GPU,且通过ChainRules提供自动微分支持。

批量矩阵-矩阵乘积

using TensorOperations, CUDA, Zygote

# CPU示例
A_cpu = rand(10, 20, 5)  # 形状(m, n, N_a)
B_cpu = rand(20, 10, 3)  # 形状(n, p, N_b)
@tensor C_cpu[im, ip, ia, ib] := A_cpu[im, in, ia] * B_cpu[in, ip, ib]

# GPU示例(直接转CuArray即可)
A_gpu = CuArray(A_cpu)
B_gpu = CuArray(B_cpu)
@tensor C_gpu[im, ip, ia, ib] := A_gpu[im, in, ia] * B_gpu[in, ip, ib]

# 梯度测试(Zygote自动兼容)
grad_A, grad_B = gradient((A, B) -> sum(@tensor C[im, ip, ia, ib] := A[im, in, ia] * B[in, ip, ib]), A_cpu, B_cpu)

批量矩阵-向量乘积

将向量B视为(n×N_b)的二维张量,直接通过张量收缩实现:

# CPU示例
A_cpu = rand(10, 20, 5)  # 形状(m, n, N_a)
B_cpu = rand(20, 3)       # 形状(n, N_b)
@tensor D_cpu[im, ia, ib] := A_cpu[im, in, ia] * B_cpu[in, ib]

# GPU示例
A_gpu = CuArray(A_cpu)
B_gpu = CuArray(B_cpu)
@tensor D_gpu[im, ia, ib] := A_gpu[im, in, ia] * B_gpu[in, ib]

方案2:NNlib.jl(高性能)

利用NNlib的batched_gemm结合维度扩展(无冗余内存分配),底层调用优化的BLAS/CUBLAS实现,性能更优。

批量矩阵-矩阵乘积

通过扩展维度将批次索引对齐:把A扩展为(m×n×N_a×1),B扩展为(n×p×1×N_b),再指定batch_dims参数完成批量乘积:

using NNlib, CUDA, Zygote

# CPU示例
A_cpu = rand(10, 20, 5)
B_cpu = rand(20, 10, 3)
# 自动广播对齐批次维度,输出形状(m, p, N_a, N_b)
C_cpu = batched_gemm(A_cpu, B_cpu; batch_dims=(3,4))

# GPU示例
A_gpu = CuArray(A_cpu)
B_gpu = CuArray(B_cpu)
C_gpu = batched_gemm(A_gpu, B_gpu; batch_dims=(3,4))

批量矩阵-向量乘积

将向量B重塑为(n×1×N_b)的张量(视为列矩阵),调用batched_gemm后再挤压冗余维度:

# CPU示例
A_cpu = rand(10, 20, 5)
B_cpu = rand(20, 3)
# 重塑B为(n,1,N_b),完成批量乘积后去掉维度1,输出形状(m, N_a, N_b)
D_cpu = squeeze(batched_gemm(A_cpu, reshape(B_cpu, 20, 1, 3); batch_dims=(3,3)), dims=2)

# GPU示例
A_gpu = CuArray(A_cpu)
B_gpu = CuArray(B_cpu)
D_gpu = squeeze(batched_gemm(A_gpu, reshape(B_gpu, 20, 1, 3); batch_dims=(3,3)), dims=2)

方案对比

  • TensorOperations.jl:语法贴合数学定义,无需手动处理维度扩展,适合复杂张量运算场景,代码可读性更高。
  • NNlib.jl:底层调用硬件优化的线性代数库,性能表现更优,适合对性能要求极高的场景,仅需少量维度调整操作。

两种方案均无需手动定义梯度,且完美兼容CPU与GPU环境。

内容的提问来源于stack exchange,提问作者max xilian

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 18:43:18