PyTorch中两次掩码操作赋值失败,如何实现预期效果?
PyTorch掩码连续索引赋值失效的解决方法
问题原因
执行result[mask1][mask2] = values[mask1][mask2]时,第一次索引result[mask1]返回的是原张量的副本而非视图。后续对这个副本的[mask2]赋值操作只会修改临时生成的副本,不会同步到原result张量,因此赋值未生效。
解决方法
方法1:合并掩码为原张量维度的完整掩码
生成一个和result同维度的最终掩码,直接对原张量赋值:
import torch values = torch.tensor([0, 0.5, 0.99, 0.87]) saved_values = values + torch.tensor([0.1, -0.4, 0, 0.1]) result = torch.zeros_like(values) mask1 = values > 0 # 创建与原张量同维度的空掩码 mask2_full = torch.zeros_like(mask1, dtype=torch.bool) # 将mask2的结果填充到mask1对应的位置 mask2_full[mask1] = ~torch.greater(saved_values[mask1], values[mask1]) # 直接使用合并后的掩码赋值 result[mask2_full] = values[mask2_full]
测试验证:
>>> result tensor([0.0000, 0.5000, 0.9900, 0.0000]) >>> result[mask2_full] tensor([0.5000, 0.9900])
方法2:使用torch.where直接完成条件赋值
通过torch.where一次性完成条件判断与赋值,无需分步处理掩码:
result = torch.zeros_like(values) # 组合条件:满足mask1 且 saved_values <= values condition = mask1 & (~torch.greater(saved_values, values)) result = torch.where(condition, values, result)
方法3:通过索引定位直接赋值
先获取满足条件的原张量索引,再直接赋值:
result = torch.zeros_like(values) mask1 = values > 0 mask2 = ~torch.greater(saved_values[mask1], values[mask1]) # 获取mask1对应的原张量索引 indices_mask1 = torch.nonzero(mask1).squeeze() # 筛选出最终需要赋值的索引 final_indices = indices_mask1[mask2] # 直接对原张量的指定索引赋值 result[final_indices] = values[final_indices]
内容的提问来源于stack exchange,提问作者MaKaNu
相关产品推荐
相关产品推荐

