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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 13:51:14