PyTorch张量指定索引与条件下的元素修改问题求助
多列场景下PyTorch张量行内1的保留处理方案
已知输入
原始张量
import torch x = torch.tensor([[1, 0, 0, 0, 0, 0, 0], [0, 1, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [1, 0, 0, 1, 0, 0, 0], [0, 0, 0, 0, 1, 1, 0], [0, 0, 1, 0, 0, 1, 0], [0, 0, 0, 1, 0, 0, 1], [0, 0, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [1, 0, 0, 0, 0, 0, 0]])
目标列索引
col_inds = torch.tensor([3,4,5])
核心需求
- 先定位出
col_inds指定列中包含1的行 - 把这些行的所有元素置0,仅保留该行在
col_inds里第一个出现的1对应的位置为1 - 最终每行仅留一个1(比如索引为4的行,
col_inds里第一个1在列4,所以要把列5的1置0)
解决方案
步骤1:高效定位目标行与保留列
用张量批量操作替代循环,效率更高:
# 获取在指定列中存在1的行索引(用any避免重复索引,比原方法更高效) row_inds = torch.where(x[:, col_inds].any(dim=1))[0] # 提取目标行在指定列的子张量 target_sub = x[row_inds][:, col_inds] # 找到每行第一个1在子张量中的列位置 first_one_pos = target_sub.argmax(dim=1) # 映射回原张量的列索引 keep_col = col_inds[first_one_pos]
步骤2:批量重置目标行
# 把目标行全部置0 x[row_inds] = 0 # 在指定位置恢复1 x[row_inds, keep_col] = 1
完整运行代码
import torch # 原始张量 x = torch.tensor([[1, 0, 0, 0, 0, 0, 0], [0, 1, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [1, 0, 0, 1, 0, 0, 0], [0, 0, 0, 0, 1, 1, 0], [0, 0, 1, 0, 0, 1, 0], [0, 0, 0, 1, 0, 0, 1], [0, 0, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [1, 0, 0, 0, 0, 0, 0]]) col_inds = torch.tensor([3,4,5]) # 获取目标行索引 row_inds = torch.where(x[:, col_inds].any(dim=1))[0] # 定位要保留的列 target_sub = x[row_inds][:, col_inds] first_one_pos = target_sub.argmax(dim=1) keep_col = col_inds[first_one_pos] # 重置行并保留指定位置的1 x[row_inds] = 0 x[row_inds, keep_col] = 1 print(x)
输出结果
tensor([[1, 0, 0, 0, 0, 0, 0], [0, 1, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [0, 0, 0, 1, 0, 0, 0], [0, 0, 0, 0, 1, 0, 0], [0, 0, 0, 0, 0, 1, 0], [0, 0, 0, 1, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0], [0, 0, 1, 0, 0, 0, 0], [1, 0, 0, 0, 0, 0, 0]])
内容的提问来源于stack exchange,提问作者user1340852
相关产品推荐
相关产品推荐

