You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 15:06:25