使用tri... 函数集时,检查源代码会很有用。它们都是python,并且基于np.tri。
制作一个小样本数组 - 说明并验证答案:
In [205]: arr = np.arange(18).reshape(2,3,3) # arange(1,19) might be better
In [206]: arr
Out[206]:
array([[[ 0, 1, 2],
[ 3, 4, 5],
[ 6, 7, 8]],
[[ 9, 10, 11],
[12, 13, 14],
[15, 16, 17]]])
tril 将上三角形的值设置为 0。它在这种情况下有效,但没有记录到 3d 数组的应用。
In [207]: np.tril(arr)
Out[207]:
array([[[ 0, 0, 0],
[ 3, 4, 0],
[ 6, 7, 8]],
[[ 9, 0, 0],
[12, 13, 0],
[15, 16, 17]]])
但在代码中 if first 从最后 2 个维度构造一个布尔掩码:
In [208]: mask = np.tri(*arr.shape[-2:], dtype=bool)
In [209]: mask
Out[209]:
array([[ True, False, False],
[ True, True, False],
[ True, True, True]])
并使用np.where 将一些值设置为0。这在3d 情况下通过广播起作用。 mask 和 arr 匹配最后 2 个维度,所以 mask 可以匹配 broadcast:
In [210]: np.where(mask, arr, 0)
Out[210]:
array([[[ 0, 0, 0],
[ 3, 4, 0],
[ 6, 7, 8]],
[[ 9, 0, 0],
[12, 13, 0],
[15, 16, 17]]])
您的tril_indices 只是这个掩码的索引:
In [217]: np.nonzero(mask) # aka np.where
Out[217]: (array([0, 1, 1, 2, 2, 2]), array([0, 0, 1, 0, 1, 2]))
In [218]: np.tril_indices(3)
Out[218]: (array([0, 1, 1, 2, 2, 2]), array([0, 0, 1, 0, 1, 2]))
它们不能直接用于索引arr:
In [220]: arr[np.tril_indices(3)].shape
Traceback (most recent call last):
File "<ipython-input-220-e26dc1f514cc>", line 1, in <module>
arr[np.tril_indices(3)].shape
IndexError: index 2 is out of bounds for axis 0 with size 2
In [221]: arr[:,np.tril_indices(3)].shape
Out[221]: (2, 2, 6, 3)
但是解压两个索引数组:
In [222]: I,J = np.tril_indices(3)
In [223]: I,J
Out[223]: (array([0, 1, 1, 2, 2, 2]), array([0, 0, 1, 0, 1, 2]))
In [224]: arr[:,I,J]
Out[224]:
array([[ 0, 3, 4, 6, 7, 8],
[ 9, 12, 13, 15, 16, 17]])
布尔掩码也可以直接使用:
In [226]: arr[:,mask]
Out[226]:
array([[ 0, 3, 4, 6, 7, 8],
[ 9, 12, 13, 15, 16, 17]])
基础 np.tri 只需在索引上执行外部 >= 即可
In [231]: m = np.greater_equal.outer(np.arange(3),np.arange(3))
In [232]: m
Out[232]:
array([[ True, False, False],
[ True, True, False],
[ True, True, True]])
In [234]: np.arange(3)[:,None]>=np.arange(3)
Out[234]:
array([[ True, False, False],
[ True, True, False],
[ True, True, True]])