如何在PyTorch中为张量每一行随机设置数量可变的1元素,满足各行上限要求
PyTorch实现每行1数量上限不同的随机0-1矩阵
以下是完全向量化的实现方案,无Python侧循环,支持CPU/GPU运行,可满足你生成n×n随机0-1矩阵、每行1的个数不超过对应行指定上限的需求。
实现代码
版本1:每行1的个数刚好等于指定上限
import torch n = 5 # 矩阵维度 row_max_ones = [2, 3, 1, 4, 2] # 长度为n的列表,对应每行1的数量上限 # 1. 初始化全0矩阵 binary_mat = torch.zeros((n, n), dtype=torch.int8) # 2. 生成每行的随机排列索引,相当于对每行n个位置做无放回打乱 rand_indices = torch.rand(n, n).argsort(dim=1) # 3. 生成掩码:每行前k个位置为True,k为对应行的1的数量 k_tensor = torch.tensor(row_max_ones) pos_mask = torch.arange(n)[None, :] < k_tensor[:, None] # 4. 按照随机索引把掩码对应的位置设为1 binary_mat.scatter_(dim=1, index=rand_indices, src=pos_mask.to(binary_mat.dtype))
版本2:每行1的个数不超过指定上限(0到上限之间随机取值)
只需要在上述代码基础上,把固定的上限替换为随机生成的实际1的数量即可:
# 替换版本1中的k_tensor定义即可 k_tensor = torch.randint(low=0, high=torch.tensor(row_max_ones) + 1, size=(n,))
方案说明
- 核心逻辑是通过打乱每行的索引来实现1的随机分布,避免了逐行循环的性能损耗,大尺寸矩阵下的运行效率远高于Python循环实现
scatter_是PyTorch的原位操作,直接在原张量上修改值,无需额外生成大尺寸中间张量,内存开销低
内容的提问来源于stack exchange,提问作者Ru11
相关产品推荐
相关产品推荐

