批次索引不相等的批量矩阵与向量乘积实现方案问询
高效实现跨批次的批量矩阵/矩阵-向量乘积(支持CPU/GPU)
问题背景
需要实现两种跨批次的张量运算,要求同时支持CPU和GPU,且避免冗余内存分配:
- 批量矩阵-矩阵乘积:给定形状为
(m×n×N_a)的张量A和(n×p×N_b)的张量B,输出C[:,:,ia,ib] = A[:,:,ia] * B[:,:,ib](最终形状(m×p×N_a×N_b)) - 批量矩阵-向量乘积:给定形状为
(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
相关产品推荐
相关产品推荐

