如何从numpy二进制矩阵中随机无重复抽取t个值为1的索引对
numpy二进制数组抽取指定数量值为1的不重复索引方案
实现思路
- 先用
np.where()提取数组中所有值为1的元素的行、列索引 - 拼接索引得到所有符合条件的坐标对,按需调整索引为1基(匹配你给出的示例输出格式)
- 对坐标对做不放回随机抽样,选取
t个不重复的索引对转为列表即可
完整实现代码
import numpy as np def sample_one_indices(A: np.ndarray, t: int, base_one: bool = True) -> list[tuple]: # 提取所有值为1的索引 row, col = np.where(A == 1) # 校验1的数量是否足够抽取 if len(row) < t: raise ValueError(f"数组中值为1的元素共{len(row)}个,少于要求抽取的{t}个") # 拼接为坐标数组 coords = np.column_stack((row, col)) # 调整为1基索引(默认开启,匹配示例输出) if base_one: coords += 1 # 不放回随机抽样t个样本 sampled_idx = np.random.choice(len(coords), size=t, replace=False) # 转为元组列表返回 return [tuple(coord) for coord in coords[sampled_idx]] # 测试示例 if __name__ == "__main__": A = np.array([[0, 1, 1], [1, 0, 1], [1, 1, 0]]) t = 2 # 可选:设置随机种子固定输出结果,方便复现 # np.random.seed(42) print(sample_one_indices(A, t)) # 输出示例:[(2, 1), (3, 2)]
补充说明
- 若你需要使用numpy默认的0基索引,调用函数时传入
base_one=False即可 - 内置样本量校验逻辑,避免1的数量不足时程序异常
内容的提问来源于stack exchange,提问作者J.Maisel
相关产品推荐
相关产品推荐

