Python numpy生成指定规则四元组索引数组的方法
实现方案
你要的四元组本质是arr1自笛卡尔积、arr2自笛卡尔积再做笛卡尔积,顺序为第一列变化最慢、第四列变化最快,和你给出的示例顺序完全匹配,用numpy内置函数可以零循环实现,性能完全满足大n下的张量索引需求。
用meshgrid实现(最直观不易错)
np.meshgrid可以直接生成多维度的笛卡尔积网格,指定indexing='ij'即可匹配矩阵维度顺序,展平后拼接就是目标结果:
import numpy as np n = 2 arr1 = np.array([0, 1]) arr2 = np.array([3, 4]) # 生成四个维度的坐标网格,顺序对应四元组的四列 col1, col2, col3, col4 = np.meshgrid(arr1, arr1, arr2, arr2, indexing='ij') quads = np.column_stack((col1.ravel(), col2.ravel(), col3.ravel(), col4.ravel()))
运行后得到的quads和你给出的示例输出完全一致,总长度为n**4,且对arr1、arr2的取值无要求,不需要是连续整数。
原tile/repeat写法的错误修正
你之前的写法错误在于重复次数和tile次数的计算不符合维度变化频率,正确的重复规则为:
- 第1列:变化最慢,每个元素重复
n**3次,无需tile - 第2列:每个元素重复
n**2次,整体tile n次 - 第3列:每个元素重复n次,整体tile
n**2次 - 第4列:变化最快,每个元素重复1次,整体tile
n**3次
对应代码:
col1 = np.repeat(arr1, n**3) col2 = np.tile(np.repeat(arr1, n**2), n) col3 = np.tile(np.repeat(arr2, n), n**2) col4 = np.tile(arr2, n**3) quads_rep = np.column_stack((col1, col2, col3, col4))
该写法结果和meshgrid版本完全一致,但重复次数计算容易出错,更推荐用meshgrid实现。
张量索引使用
生成的四元组用于索引4维张量时,拆分为四个一维数组传入索引效率最高,避免numpy对二维索引做额外处理:
# 示例4维张量 T = np.random.rand(n, n, n, n) # 提取对应位置的值 result = T[quads[:,0], quads[:,1], quads[:,2], quads[:,3]]
两种实现都是纯numpy向量化操作,无Python层循环,性能远高于itertools.product或列表推导的实现,n较大时优势明显。
内容的提问来源于stack exchange,提问作者imlg
相关产品推荐
相关产品推荐

