获取a[i][j] == 1的索引
您可以使用numpy.argwhere() 或numpy.nonzero() 有效地获取您正在寻找的数据(0 和1 矩阵中1 的位置),但是您将无法以中指定的格式获取它们您的原始问题仅使用 NumPy ndarrays。
您可以使用 ndarrays 和标准 Python 列表的组合来获得指定格式的数据,但是鉴于您正在使用的数据的大小,效率是最重要的,我认为最好专注于获取数据而不是以不规则 Python 列表的 ndarray 格式获取它。
如果您提到的格式是硬性要求,您始终可以在计算后重新格式化结果(矩阵中 1 的索引),这样您的代码将受益于 NumPy 在繁重计算期间提供的优化 - 减少整个过程的执行时间。
使用np.argwhere()的示例
import numpy as np
a = np.random.randint(0, 2, size=(4,4))
b = np.argwhere(a == 1)
print(f'a\n{a}')
print(f'b\n{b}')
输出
a
[[1 1 1 1]
[0 0 0 0]
[1 0 1 0]
[1 1 1 1]]
b
[[0 0]
[0 1]
[0 2]
[0 3]
[2 0]
[2 2]
[3 0]
[3 1]
[3 2]
[3 3]]
如您所见,np.argwhere(a == 1) 返回一个 ndarray,其值是包含 a 中位置索引的 ndarray,其值 (x) 满足条件 x == 1。
我用a = np.random.randint(0, 2, size=(10000,10000) 在我的笔记本电脑上尝试了几次上述方法(没什么花哨的),每次大约 3-5 秒完成。
获取所有值!= 1的行索引
如果您想存储 a 的所有行索引不包含值 == 1,最直接的方法(假设您使用我上面的示例代码)可能是使用 numpy.setdiff1d() 返回一个行数组b 中不存在的索引 - 即包含 a 的所有行索引的数组与一维数组 b[0] 之间的设置差异,这将是 a 中所有值的行索引 != 1。
假设a 和b 与上例相同。
c = np.setdiff1d(np.arange(a.shape[0]), b[:, 0])
print(c)
输出
array([1])
在上面的示例中,c = [1] 1 是 a 中唯一不包含任何值 == 1 的行索引。
值得注意的是,如果a 定义为np.random.randint(0, 2, size=(10000,10000),则c 不是零长度(即空)数组的概率非常小。这是因为如果一行不包含值 == 1,np.random 必须连续返回 0 10,000 次才能用 0 填充一行。
为什么要使用多个 NumPy 数组?
我知道使用b 和c 分别存储与a == 1 和a != 1 所在位置相关的结果可能看起来很奇怪。为什么不使用原始问题中所述的不规则 list?
简而言之,就是效率。通过使用 NumPy 数组,您将能够对数据进行矢量化计算,并在很大程度上避免代价高昂的 Python 循环,其好处将被显着放大,这反映在考虑到您正在处理的数据大小的执行时间上。
您始终可以以更人性化的不同格式存储数据,并根据需要将其映射回 NumPy,但是与原始问题中的示例相比,上述示例可能会在执行时显着提高效率。