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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 01:36:05