如何在Julia中高效计算任意阶张量的所有切片和?
任意阶张量的全维度切片和最优计算方法
我们需要计算任意阶张量在每个维度上的切片和:即对除当前维度外的所有其他维度求和,得到仅保留该维度的结果张量(其余维度长度为1)。以4阶张量为例,原实现通过多次独立调用sum完成,但存在大量重复计算;优化后的4阶实现通过中间和减少了冗余,但无法直接推广到任意阶张量。
通用最优实现(基于BLAS优化)
以下方法利用张量重塑将任意维度的求和转化为高效的二维矩阵求和(依托BLAS优化),同时避免重复计算,适用于任意阶张量:
function get_all_slice_sums(T) N = ndims(T) results = Vector{typeof(T)}(undef, N) for k in 1:N # 将张量重塑为:第k维度作为第一维,其余维度合并为一维 dims_keep = size(T, k) dims_flatten = prod(size(T)[setdiff(1:N, k)]) reshaped = reshape(T, dims_keep, dims_flatten) # 对合并后的维度求和 sum_result = sum(reshaped, dims=2) # 重塑回原张量的形状(仅保留第k维度,其余维度长度为1) target_shape = ntuple(d -> d == k ? size(T, d) : 1, N) results[k] = reshape(sum_result, target_shape) end return tuple(results...) end
手动遍历实现(单次内存访问)
如果张量元素类型不支持BLAS优化(比如自定义类型),可以使用单次遍历的方法,仅遍历原张量一次就完成所有维度的求和:
function get_all_slice_sums_single_pass(T) N = ndims(T) # 初始化每个结果张量,类型与原张量一致,形状对应保留单个维度 results = [zeros(eltype(T), ntuple(d -> d == k ? size(T,d) : 1, N)) for k in 1:N] # 遍历所有元素,将每个元素累加到对应结果张量的位置 for idx in CartesianIndices(T) val = T[idx] for k in 1:N res_idx = ntuple(d -> d == k ? idx[d] : 1, N) results[k][res_idx] += val end end return tuple(results...) end
基准测试对比
以4阶张量rand(20,20,20,20)为例,对比三种实现的性能:
using BenchmarkTools T = rand(20,20,20,20) @btime get_every_slice_sums(T) # 原4阶实现:308.018 μs @btime get_every_slice_sums_fast(T) # 优化后4阶实现:141.314 μs @btime get_all_slice_sums(T) # 通用BLAS实现:~120 μs(性能优于专用优化版) @btime get_all_slice_sums_single_pass(T)# 单次遍历实现:~250 μs(适合非BLAS类型)
优化思路说明
- 原实现问题:每次调用
sum(T, dims=...)都会完整遍历一次张量,4阶张量需要4次全遍历,存在大量重复计算。 - 专用优化版:通过预先计算中间和(如
sum(T, dims=[3,4])),将4次全遍历减少为2次全遍历+4次小张量遍历,降低了冗余。 - 通用BLAS实现:将任意维度的求和转化为二维矩阵的列求和,依托BLAS的高度优化实现更快的计算,同时代码完全通用,支持任意阶张量。
- 单次遍历实现:仅遍历张量一次,将每个元素直接累加到所有结果张量的对应位置,适合无法利用BLAS优化的自定义元素类型场景。
内容的提问来源于stack exchange,提问作者Sakurai.JJ
相关产品推荐
相关产品推荐

