如何在NumPy中对多维索引执行外减并调整轴生成任意维度距离表
任意维度空间点欧氏距离查找表实现问题
我需要计算空间中点的欧氏(或其他)距离查找表,用于空间内的距离查询。在二维场景下,可通过table[x1,y1,x2,y2]进行查询,例如:
table[0,0,0,10]应返回10table[0,0,10,10]应返回14.142135623730951
现有仅支持二维场景的代码
import numpy as np to = np.indices((5, 5)).T x_shape, y_shape, n_dim = to.shape x_sub = np.subtract.outer(to[..., 0], to[..., 0]) y_sub = np.subtract.outer(to[..., 1], to[..., 1]) distances = np.sqrt(x_sub ** 2 + y_sub ** 2) n_test = 5 for i in np.random.choice(x_shape, n_test): for j in np.random.choice(y_shape, n_test): for k in np.random.choice(x_shape, n_test): for l in np.random.choice(y_shape, n_test): d = np.sqrt((i - k) ** 2 + (j - l) ** 2) assert (distances[i, j, k, l] == d) print('\033[92m TEST PASSES')
问题描述
我希望将上述代码改写为支持任意维度的版本。虽然可以通过循环对每个维度执行np.subtract.outer来实现,但尝试直接对索引执行外减操作(如np.subtract.outer(to, to))时,得到的结果形状为(5, 5, 2, 5, 5, 2),而我需要的形状是(5, 5, 5, 5, 2),请问该如何调整实现?
解决方案代码
import numpy as np to = np.indices((5, 5, 5)).T sh = list(to.shape) n_dims = sh[-1] t = to.reshape(sh[:-1] + [1] * n_dims + [n_dims]) - to.reshape([1] * n_dims + sh[:-1] + [n_dims]) distances = np.sqrt(np.sum(np.power(t, 2), axis=-1)) n_test = 100 idxs = np.indices(distances.shape).T.reshape(-1, n_dims * 2) for idx in np.random.permutation(idxs)[:n_test]: assert distances[tuple(idx)] == np.linalg.norm(idx[:n_dims] - idx[n_dims:]) print('\033[92m TEST PASSES')
内容的提问来源于Stack Exchange,提问作者imkded5
相关产品推荐
相关产品推荐

