如何提取对称张量的排列唯一元素及对应索引?
提取对称张量的排列唯一元素及索引
要解决这个问题,核心是基于索引的排列等价性分组,而非依赖元素值去重。np.unique会误将偶然相等但索引不排列等价的元素归为一类,因此我们需要从索引本身的结构入手处理:
实现思路
- 生成张量的所有索引集合
- 为每个索引生成「规范标识」:将索引元组排序,排列等价的索引排序后完全相同
- 基于规范标识去重,保留每个等价类的一个原始索引及对应元素
代码示例
import numpy as np def get_perm_unique_elements(tensor): # 获取张量的阶数n和单维度大小m n = tensor.ndim m = tensor.shape[0] # 生成所有索引:先得到(n, m, m, ..., m)的索引网格,转置后重塑为(m^n, n)的二维数组 indices = np.indices(tensor.shape).transpose().reshape(-1, n) # 对每个索引元组排序,得到排列等价类的统一标识 sorted_indices = np.sort(indices, axis=1) # 提取唯一规范标识对应的原始索引位置 _, unique_positions = np.unique(sorted_indices, axis=0, return_index=True) # 整理唯一索引和对应元素 unique_idx_tuples = [tuple(idx) for idx in indices[unique_positions]] unique_values = tensor[tuple(indices[unique_positions].T)] return list(zip(unique_idx_tuples, unique_values)) # 测试:构造一个3阶对称张量(shape=(2,2,2)) m, n = 2, 3 # 先生成排列唯一元素的基础值,再填充所有排列等价的索引位置 base_values = np.random.randint(0, 10, size=6) sym_tensor = np.zeros((m,)*n) for idx in np.ndindex(sym_tensor.shape): sorted_idx = tuple(sorted(idx)) # 映射排序后的索引到基础值的位置(简单哈希实现) val_pos = sum([m**i * sorted_idx[i] for i in range(n)]) - sum([m**i * i for i in range(n)]) sym_tensor[idx] = base_values[val_pos] # 提取结果 unique_items = get_perm_unique_elements(sym_tensor) print("排列唯一索引及对应元素:") for idx, val in unique_items: print(f"索引{idx}: {val}")
关键说明
- 索引规范化:排序索引元组是最直接的方式,确保所有排列等价的索引拥有完全相同的标识,彻底避免元素值偶然相等带来的误判。
- 高效去重:
np.unique(axis=0)对排序后的索引数组按行去重,return_index=True直接获取原始索引的位置,无需额外循环分组,效率较高。 - 通用性:代码适用于任意阶数n和任意维度m的对称张量,无需针对特定阶数修改逻辑。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

