PyTorch中掩码聚合(均值、最值等)是否有内置实现方法?
PyTorch中实现掩码聚合运算的内置方法/工具包
示例张量
import torch x = torch.tensor([ [1, 2, -1, -1], [10, 20, 30, -1] ]) mask = torch.tensor([ [True, True, False, False], [True, True, True, False] ])
手动实现的掩码均值
你已经手动实现了掩码均值的计算逻辑:
n_mask = torch.sum(mask, axis=1) x_mean = torch.sum(x * mask, axis=1) / n_mask print(x_mean)
输出结果:
tensor([ 1.50, 20.00])
内置方案实现各类掩码聚合
PyTorch提供了多种内置方式,可以更简洁地完成掩码均值、最大值、最小值等聚合操作:
掩码均值
利用torch.nanmean自动忽略nan的特性,先将掩码外的元素替换为nan:
# 整数张量无法存储nan,需先转为浮点型 x_masked = torch.where(mask, x.float(), torch.tensor(torch.nan)) x_mean = torch.nanmean(x_masked, dim=1) print(x_mean) # 输出: tensor([1.50, 20.00])
掩码最大值
将掩码外的元素替换为负无穷,这样取最大值时会自动忽略这些无效值:
x_masked_max = torch.where(mask, x.float(), torch.tensor(-torch.inf)) x_max = torch.max(x_masked_max, dim=1)[0] print(x_max) # 输出: tensor([2., 30.])
掩码最小值
将掩码外的元素替换为正无穷,取最小值时会自动忽略无效值:
x_masked_min = torch.where(mask, x.float(), torch.tensor(torch.inf)) x_min = torch.min(x_masked_min, dim=1)[0] print(x_min) # 输出: tensor([1., 10.])
掩码求和
除了手动计算torch.sum(x * mask, dim=1),也可以用torch.where替换无效值后求和,效果一致:
x_sum = torch.sum(torch.where(mask, x, torch.tensor(0)), dim=1) print(x_sum) # 输出: tensor([ 3, 60])
如果使用PyTorch生态的高阶工具包(如TorchText、TorchVision),部分场景会封装好掩码聚合工具,但基础场景用上述内置方法完全足够。
内容的提问来源于stack exchange,提问作者Sengiley
相关产品推荐
相关产品推荐

