如何利用NumPy索引加速从降维张量恢复对称全张量?
问题
我有一组对称张量full_tensors,已将其中所有不同元素提取到二维数组tensor_reduced中。现在需要从tensor_reduced还原出完整的full_tensors。我有列表all_permutations_list,其中每个元素是一组索引元组,这些索引对应的位置应被赋予相同的值(例如all_permutations_list[1]=((0, 0, 1, 0, 0, 0), (0, 1, 0, 0, 0, 0), (1, 0, 0, 0, 0, 0), (0, 0, 0, 1, 0, 0), (0, 0, 0, 0, 1, 0), (0, 0, 0, 0, 0, 1)),且tensor_reduced.shape[1]与len(all_permutations_list)相等。
当前实现代码如下:
def makeFulltensor(tensor_reduced,all_permutations_list,dim=4): full_tensors=np.zeros((tensor_reduced.shape[0],dim,dim,dim,dim,dim,dim),dtype=tensor_reduced.dtype) for index,all_permutations in enumerate(all_permutations_list): for perm in (all_permutations): full_tensors[:,perm[0],perm[1],perm[2],perm[3],perm[4],perm[5]]=tensor_reduced[:,index] return full_tensors
希望利用NumPy的索引特性加速代码,去除内层甚至外层循环,求高效实现方法。
附all_permutations_list生成代码:
from itertools import permutations from itertools import combinations_with_replacement as cwr all_permutations_list=[] reduced_indices=list(cwr(range(4), 6)) for index,(i,j,k,l,m,n) in enumerate(reduced_indices): all_permutations=set(permutations([i,j,k,l,m,n])) all_permutations_list.append(all_permutations) all_permutations_list=[tuple(x) for x in all_permutations_list]
高效实现方案
核心思路是把all_permutations_list预处理成批量索引数组,利用NumPy的高级索引一次性完成赋值,彻底消除循环。
步骤1:预处理索引,生成批量索引数组
先将all_permutations_list中的所有索引元组整理成二维数组,同时记录每个索引对应的tensor_reduced列索引:
import numpy as np from itertools import permutations, combinations_with_replacement as cwr # 重新生成并预处理索引(如果已有all_permutations_list,也可直接基于它处理) reduced_indices = list(cwr(range(4), 6)) idxs = [] col_indices = [] for col_idx, perm_base in enumerate(reduced_indices): perms = set(permutations(perm_base)) for perm in perms: idxs.append(perm) col_indices.append(col_idx) # 转换为NumPy数组,适配高级索引 idxs = np.array(idxs) # shape: (N, 6),N为所有对称位置的总数 col_indices = np.array(col_indices) # shape: (N,)
步骤2:利用广播和高级索引完成赋值
通过一次索引操作完成所有位置的赋值,无需循环:
def makeFulltensor_fast(tensor_reduced, idxs, col_indices, dim=4): batch_size = tensor_reduced.shape[0] # 初始化全零张量 full_tensors = np.zeros((batch_size, dim, dim, dim, dim, dim, dim), dtype=tensor_reduced.dtype) # 高级索引+广播:一次性完成所有对称位置的赋值 full_tensors[:, idxs[:,0], idxs[:,1], idxs[:,2], idxs[:,3], idxs[:,4], idxs[:,5]] = tensor_reduced[:, col_indices] return full_tensors
步骤3:使用示例
# 构造测试用tensor_reduced(shape对应cwr(4,6)的长度84) tensor_reduced = np.random.rand(10, 84) # 生成预处理的索引数组(仅需生成一次,可缓存复用) reduced_indices = list(cwr(range(4), 6)) idxs = [] col_indices = [] for col_idx, perm_base in enumerate(reduced_indices): perms = set(permutations(perm_base)) for p in perms: idxs.append(p) col_indices.append(col_idx) idxs = np.array(idxs) col_indices = np.array(col_indices) # 调用快速函数 full_tensors_fast = makeFulltensor_fast(tensor_reduced, idxs, col_indices) # 验证结果与原函数一致 full_tensors_original = makeFulltensor(tensor_reduced, all_permutations_list) print(np.allclose(full_tensors_fast, full_tensors_original)) # 输出True表示结果正确
性能说明
- 原代码是双重循环,时间复杂度为O(M*K)(M为
all_permutations_list长度,K为每组排列数); - 优化后利用NumPy的C级批量索引操作,速度可提升几十到上百倍(取决于张量规模);
- 可提前缓存与
dim对应的索引数组,避免重复生成,进一步提升效率。
内容的提问来源于stack exchange,提问作者cheetah
相关产品推荐
相关产品推荐

