PyTorch多层掩码下Tensor赋值失效问题及简便方案咨询
解决PyTorch中连续掩码赋值不生效的简洁方法
你判断的原因完全正确:x[mask]返回的是原张量的副本而非视图,所以x[mask][second_mask] = 100这种连续索引赋值只会修改临时副本,原张量不会产生任何变化。
下面是两种更简洁的实现方式,无需创建临时张量再回写:
方法1:合并布尔掩码
先基于原掩码生成最终的布尔掩码,直接对原张量赋值:
import torch x = torch.linspace(1,9,9).reshape((3,3)) mask = x > 5 # 实际场景中的布尔型second_mask示例 second_mask = torch.tensor([True, False, True]) # 生成最终掩码:仅保留mask中被second_mask选中的位置 final_mask = mask.clone() final_mask[mask] = second_mask # 直接赋值 x[final_mask] = 100 print(x)
方法2:通过索引筛选赋值
先提取原掩码对应的非零索引,再用second_mask筛选出目标索引,直接对原张量的指定位置赋值:
import torch x = torch.linspace(1,9,9).reshape((3,3)) mask = x > 5 second_mask = torch.tensor([True, False, True]) # 获取mask对应的所有索引(元组形式,对应每个维度的索引) mask_indices = torch.nonzero(mask, as_tuple=True) # 用second_mask筛选出需要修改的索引 target_indices = tuple(idx[second_mask] for idx in mask_indices) # 直接赋值 x[target_indices] = 100 print(x)
两种方法都能直接修改原张量,避免了临时张量的来回赋值操作,其中方法1更适合布尔掩码的场景,代码可读性更高;方法2在需要更灵活处理索引时会更实用。
内容的提问来源于stack exchange,提问作者xh c
相关产品推荐
相关产品推荐

