注意
跳转至页面底部 下载完整示例代码。
MaskedTensor 高级语义#
在学习本教程之前,请务必先查阅我们的 MaskedTensor 概述教程 <https://pytorch.ac.cn/tutorials/prototype/maskedtensor_overview.html>。
本教程旨在帮助用户理解一些高级语义的工作原理及其由来。我们将重点讨论其中两个方面:
*. MaskedTensor 与 NumPy 的 MaskedArray 之间的区别 *. 归约语义
准备工作#
# Disable prototype warnings and such
MaskedTensor 与 NumPy 的 MaskedArray#
NumPy 的 MaskedArray 在语义上与 MaskedTensor 有几处根本的区别。
- *. 它们的工厂函数和基本定义反转了掩码(类似于
torch.nn.MHA);也就是说,MaskedTensor 使用
True表示“已指定”,使用False表示“未指定”(或“有效”/“无效”),而 NumPy 则恰恰相反。我们认为,我们的掩码定义不仅更直观,而且与整个 PyTorch 中现有的语义更加一致。- *. 交集语义。在 NumPy 中,如果两个元素中的任意一个被掩码屏蔽,则结果元素也会
被屏蔽——实际上,它们 应用了 logical_or 运算符。
与此同时,MaskedTensor 不支持对掩码不匹配的操作数进行加法或二元运算——要了解原因,请参阅 归约语义部分。
然而,如果确实需要这种行为,MaskedTensor 也通过提供对数据和掩码的访问,并利用 to_tensor() 将 MaskedTensor 转换为填充了掩码值的 Tensor,从而支持此类语义。例如:
请注意,掩码为 mt0.get_mask() & mt1.get_mask(),因为 MaskedTensor 的掩码是 NumPy 掩码的逆。
归约语义#
回想一下 MaskedTensor 概述教程 中提到的“实现缺失的 torch.nan* 操作”。这些都是归约的例子——即从张量中移除一个(或多个)维度并聚合结果的运算符。在本节中,我们将使用归约语义来论证前文中提到的关于匹配掩码的严格要求。
从根本上讲,:class:`MaskedTensor` 在执行归约运算时会忽略被屏蔽(未指定)的值。举例说明:
现在,看看不同的归约操作(均在 dim=1 上):
值得注意的是,被屏蔽元素下的值不能保证具有任何特定值,尤其是当行或列被完全屏蔽时(归一化操作也是如此)。关于掩码语义的更多详细信息,可以查看此 RFC。
现在,我们可以重温这个问题:为什么我们要强制要求二元运算符的掩码必须匹配?换句话说,为什么我们不使用与 np.ma.masked_array 相同的语义?考虑以下示例:
现在,让我们尝试加法运算:
和与加法显然应该是结合律的,但使用 NumPy 的语义则不然,这肯定会让用户感到困惑。
另一方面,MaskedTensor 将直接禁止此操作,因为 mask0 != mask1。话虽如此,如果用户有需要,也有规避方法(例如,像下文所示,使用 to_tensor() 将 MaskedTensor 中未定义的元素填充为 0),但用户现在必须更明确地表达他们的意图。
结论#
在本教程中,我们了解了 MaskedTensor 与 NumPy 的 MaskedArray 背后的不同设计决策,以及归约语义。通常,MaskedTensor 的设计旨在避免歧义和令人困惑的语义(例如,我们尽量在二元运算中保持结合律),这反过来可能要求用户在编写代码时更加审慎,但我们认为这是更好的做法。如果您对此有任何想法,请 告知我们!
# %%%%%%RUNNABLE_CODE_REMOVED%%%%%%
脚本总运行时间:(0 分 0.002 秒)