PyTorch张量应用掩码后维度丢失,如何保持原维度?
PyTorch 应用掩码后保持张量原维度
我有一个2D(或更高维度)的PyTorch张量,应用同形状的二进制掩码后,输出变成了1维张量。如何在掩码操作后保留原张量的维度?
示例代码:
import torch x = torch.tensor([[1.0, 2.0, 8.0], [-4.0, 0.0, 3.0]]) mask = x >=2.0 print(x[mask]) # 输出: tensor([2., 8., 3.])
这里输出是1维,而我需要得到和原张量形状一致的2维结果。
解决方案
方法1:使用torch.where
torch.where可以根据掩码选择对应位置的值,不符合条件的位置可以指定填充值(比如0、NaN等),完美保留原张量维度:
import torch x = torch.tensor([[1.0, 2.0, 8.0], [-4.0, 0.0, 3.0]]) mask = x >=2.0 # 不符合掩码条件的位置填充0 result = torch.where(mask, x, torch.tensor(0.0, device=x.device)) print(result) # 输出: tensor([[0., 2., 8.], # [0., 0., 3.]])
如果不想用固定值填充,也可以替换成torch.nan或者其他自定义值:
# 用NaN填充不符合条件的位置 result = torch.where(mask, x, torch.nan)
方法2:使用masked_fill
张量自带的masked_fill方法也能实现需求,注意这里要传入反向掩码(~mask),指定不符合条件位置的填充值:
import torch x = torch.tensor([[1.0, 2.0, 8.0], [-4.0, 0.0, 3.0]]) mask = x >=2.0 # 对不符合掩码的位置填充0 result = x.masked_fill(~mask, 0.0) print(result) # 输出: tensor([[0., 2., 8.], # [0., 0., 3.]])
方法3:使用MaskedTensor(PyTorch 1.10+)
如果需要保留掩码信息而非直接填充值,可以用PyTorch的MaskedTensor,它会同时存储原始数据和掩码,严格保持原维度:
import torch from torch.masked import MaskedTensor x = torch.tensor([[1.0, 2.0, 8.0], [-4.0, 0.0, 3.0]]) mask = x >=2.0 mt = MaskedTensor(x, mask) print(mt) # 输出: # MaskedTensor( # data=[[1.0, 2.0, 8.0], # [-4.0, 0.0, 3.0]], # mask=[[False, True, True], # [False, False, True]] # )
后续可以通过mt.data获取原始数据,mt.mask获取掩码,进行运算时会自动遵循掩码规则。
内容的提问来源于stack exchange,提问作者Matt Pitkin
相关产品推荐
相关产品推荐

