如何用np.repeat/np.tile高效实现TxNxM数组转TxMxNxN数组?
高效实现Numpy数组维度转换:TxNxM → TxMxNxN
问题描述
给定一个维度为 TxNxM 的Numpy数组:
import numpy as np arr = np.array([[[0.83, 0.3 , 0.22], [0.17, 0.33, 0.37]], [[0. , 0.28, 0.09], [0. , 0.1 , 0.18]], [[0. , 0. , 0.08], [0. , 0. , 0.05]], [[0. , 0. , 0. ], [0. , 0. , 0. ]]])
需要将其转换为维度为 TxMxNxN 的目标数组(示例输出见问题描述),现有基于 np.stack 结合 np.repeat、reshape 的实现效率较低,需用 np.repeat 或 np.tile 实现高效转换。
高效实现方案
核心思路
观察维度变换规律:
- 原数组维度
(T, N, M)→ 先交换N和M轴,得到(T, M, N)(此操作为视图操作,不复制数据) - 在最后添加一个维度,变为
(T, M, N, 1) - 在最后一个维度重复
N次,即可得到目标维度(T, M, N, N)
方法1:使用 np.repeat
# 交换N和M轴 → 添加最后一维 → 重复N次 result = np.repeat(arr.swapaxes(1, 2)[..., np.newaxis], repeats=arr.shape[1], axis=-1) print(result.shape) # 输出:(4, 3, 2, 2)
方法2:使用 np.tile
# 交换N和M轴 → 添加最后一维 → 最后一维重复N次 result_tile = np.tile(arr.swapaxes(1,2)[..., np.newaxis], reps=(1,1,1,arr.shape[1])) print(result_tile.shape) # 输出:(4, 3, 2, 2)
额外优化:广播机制(更高效,无需内存复制)
如果不需要实际复制数据(仅需视图或可广播场景),可直接利用Numpy广播特性,避免内存开销:
broadcast_result = arr.swapaxes(1,2)[..., np.newaxis] * np.ones((1,1,1,arr.shape[1]), dtype=arr.dtype)
结果验证
可通过以下代码验证转换结果是否与目标一致:
target = np.array([[[[0.83, 0.83], [0.17, 0.17]], [[0.3 , 0.3 ], [0.33, 0.33]], [[0.22, 0.22], [0.37, 0.37]]], [[[0. , 0. ], [0. , 0. ]], [[0.28, 0.28], [0.1 , 0.1 ]], [[0.09, 0.09], [0.18, 0.18]]], [[[0. , 0. ], [0. , 0. ]], [[0. , 0. ], [0. , 0. ]], [[0.08, 0.08], [0.05, 0.05]]], [[[0. , 0. ], [0. , 0. ]], [[0. , 0. ], [0. , 0. ]], [[0. , 0. ], [0. , 0. ]]]]) print(np.allclose(result, target)) # 输出:True
内容的提问来源于stack exchange,提问作者thesecond
相关产品推荐
相关产品推荐

