ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch MaskedTensor 完全指南:用 data + mask 语义构建安全、一致的掩码张量计算

PyTorch MaskedTensor 完全指南:用 data + mask 语义构建安全、一致的掩码张量计算 PyTorch MaskedTensor 完全指南用 data mask 语义构建安全、一致的掩码张量计算【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchMaskedTensor 是 PyTorch 中一个处于原型prototype阶段的torch.Tensor子类它将「数据data」与「布尔掩码mask」绑定为一个整体让用户可以用统一、明确、可自动求导的方式表达「哪些元素参与计算、哪些元素被忽略」这一语义。本文以 docs/source/masked.md 为骨架结合 torch/masked/maskedtensor 下的源码实现与 test/test_maskedtensor.py 测试用例系统讲解 MaskedTensor 的动机、构造方式、全部已支持算子一元/二元/归约/视图选择、底层 dispatch 机制、稀疏布局支持以及原型阶段的各种限制。读完本文你将能够独立使用 MaskedTensor 处理变长张量、区分 0 与 NaN 梯度、表达稀疏算子语义并在算子缺失时正确判断应如何寻求支持。一、Motivation为什么 PyTorch 需要 MaskedTensorspecified已指定与unspecified未指定两个概念在 PyTorch 中由来已久却始终没有统一的语义。原生torch.Tensor无法妥善回答这个位置的 0 到底是真实的 0还是应该被忽略的占位符这类问题由此积累了大量 issue而 MaskedTensor 正是为解决这一系列问题而生。根据文档中的定位MaskedTensor 希望达成的目标包括承载任意掩码语义例如变长张量variable length tensors、nan*系列算子nan_to_num、nanmean等等场景都可以用统一的掩码语义来表达区分 0 梯度与 NaN 梯度掩码可以让梯度传播路径明确不再需要靠填 0或填 NaN来模拟被忽略位置支撑稀疏应用MaskedTensor 直接支持 sparse COO 与 sparse CSR 的 data 与 mask详见后文稀疏布局小节。文档明确指出MaskedTensor 的核心目标是成为 PyTorch 中 specified 与 unspecified 值的唯一事实来源source of truth让这些值成为一等公民first class citizen而非事后补丁从而进一步释放 torch.sparse 的潜力使算子更安全、更一致用户体验更直观。注意PyTorch 的 masked 张量 API 目前仍处于原型阶段未来可能发生变化。在 core.py 的__new__中每次构造 MaskedTensor 都会发出UserWarning提示这一状态。二、什么是 MaskedTensor一份 data 加一份 maskMaskedTensor 本质上是一个张量子类由两部分构成inputdata实际参与计算的数据张量mask布尔张量指明 input 中哪些条目应被包含True或忽略False。上图中给出了文档原生的对比示例上面是普通张量在max(0)过程中把所有小于 0 的值替换为 0最终全 0 矩阵的最大值仍是 0下面是 MaskedTensor所有 0 值被灰显遮罩掉剩下的元素-1、-2、-3中最大值是 -1。同一个数据、同样的算子仅因是否携带 mask 就得到完全不同的结果——这正是 MaskedTensor 允许用户系统性地忽略任何元素这一灵活性的体现。在源码层面这一结构体现在 core.py 中class MaskedTensor(torch.Tensor): staticmethod def __new__(cls, data, mask, requires_gradFalse): # data、mask 都必须是 Tensor且不能是 MaskedTensor ... # 使用相同 size 的 Tensor 作为 wrapper kwargs { device: data.device, dtype: data.dtype, layout: data.layout, requires_grad: requires_grad, dispatch_sizes_strides_policy: strides, dispatch_layout: True, } return torch.Tensor._make_wrapper_subclass(cls, data.size(), **kwargs)构造完成后_preprocess_data会分别对 data 与 mask 做clone()并存入self._masked_data/self._masked_mask两个内部字段_validate_members则会做一系列严格的合法性校验详见第四节。三、快速上手构造 MaskedTensor 的两种工厂函数官方推荐的入口在 creation.py它刻意镜像了torch.tensor与torch.as_tensor的区分工厂函数语义源码行为masked_tensor(data, mask, requires_gradFalse)保证为叶子节点相当于torch.tensor直接调用MaskedTensor(data, mask, requires_grad)as_masked_tensor(data, mask)可微构造保留 autograd 历史相当于torch.as_tensor调用MaskedTensor._from_values(data, mask)其中_from_values在 core.py 中通过一个自定义torch.autograd.Function实现forward构造 MaskedTensorbackward原样回传梯度grad_output, None因此它是可微构造函数。测试文件 test/test_maskedtensor.py 同时导入了masked_tensor与as_masked_tensor进行验证。一个最基本的构造与打印示例文档中的select示例也采用同款构造方式 from torch.masked import masked_tensor data torch.arange(12, dtypetorch.float).reshape(3, 4) mask torch.tensor([[True, False, False, True], ... [False, True, False, False], ... [True, True, True, True]]) mt masked_tensor(data, mask) mt MaskedTensor( [ 0.0000, --, --, 3.0000], [ --, 5.0000, --, --], [ 8.0000, 9.0000, 10.0000, 11.0000] )可以看到被遮罩的位置在repr中显示为--这是 core.py 中_masked_tensor_str的格式化逻辑。四、源码级剖析MaskedTensor 类内部做了什么MaskedTensor 是一个标准的 wrapper subclass核心机制值得仔细拆解4.1 合法性校验_validate_members构造时的校验规则见 core.py包括类型与布局data与mask的类型必须一致type(data) is type(mask)data的 layout 仅支持torch.strided、torch.sparse_coo、torch.sparse_csr三种且data.layout mask.layout在_preprocess_data中检查稀疏一致性对 sparse COOdata与mask的indices()必须完全一致对 sparse CSRcrow_indices()与col_indices()都必须一致dtype 约束mask必须是torch.booldata仅支持float16 / float32 / float64 / bool / int8 / int16 / int32 / int64形状约束data.dim() mask.dim()且data.size() mask.size()。4.2 核心访问接口get_data()返回被掩码的数据通过自定义autograd.Function包装backward时若梯度是 MaskedTensor 则原样返回否则用self.get_mask()包成 MaskedTensor见 core.pyget_mask()返回布尔掩码core.pyto_tensor(value)把被遮罩位置填充为指定值后转回普通 Tensor等价于data.masked_fill(~mask, value)core.pyis_sparse_coo()/is_sparse_csr()/is_sparse用于判断稀疏布局core.py。4.3 分发机制torch_function与torch_dispatchMaskedTensor 通过两张分发表驱动所有算子见 core.py 与 _ops_refs.py__torch_function__查找_MASKEDTENSOR_FUNCTION_TABLE命中则直接调用对应的实现函数否则在torch._C.DisableTorchFunctionSubclass()下执行原函数并按需_convert回包装类型__torch_dispatch__将func归一化为func.overloadpacket后查找_MASKEDTENSOR_DISPATCH_TABLE未命中时打印{func.__name__} is not implemented in __torch_dispatch__ for MaskedTensor.警告并返回NotImplemented同时提示用户到pytorch/maskedtensor的 issue 区提交最小复现样例core.py。两张分发表由五个子模块注册一元unary.py、二元binary.py、归约reductions.py、透传passthrough.py以及_ops_refs.py中的若干辅助autograd.Function如_MaskedContiguous、_MaskedToDense、_MaskedToSparse、_MaskedToSparseCsr见 _ops_refs.py。五、支持的一元算子Unary Operators一元算子指只含单个输入的算子。对 MaskedTensor 应用一元算子的规则很直接在某个索引处若数据被遮罩则继续遮罩否则正常应用算子。源码层面unary.py 的_unary_helper会把 args 中的 MaskedTensor 分别映射为_masked_mask与_masked_data对 data 执行算子fn再与 mask 一起_wrap_result回 MaskedTensorinplace 版本则通过_set_data_mask就地更新。当前支持的一元算子完整列表出自文档与 unary.py 的UNARY_NAMES一一对应abs absolute acos arccos acosh arccosh angle asin arcsin asinh arcsinh atan arctan atanh arctanh bitwise_not ceil clamp clip conj_physical cos cosh deg2rad digamma erf erfc erfinv exp exp2 expm1 fix floor frac lgamma log log10 log1p log2 logit i0 isnan nan_to_num neg negative positive pow rad2deg reciprocal round rsqrt sigmoid sign sgn signbit sin sinc sinh sqrt square tan tanh trunc支持inplace版本的一元算子为上述全部减去以下 4 个它们没有 inplace 语义angle positive signbit isnan此外 unary.py 还显式维护了UNARY_NAMES_UNSUPPORTED列表标注了已知暂不支持或语义复杂的一元算子如copysign、float_power、logical_not、hypot、ldexp、real、imag、gradient、frexp、xlogy等便于开发者定位缺口。值得一提的是一元算子在稀疏布局下同样可用_unary_helper对 sparse COO 会先coalesce()再对values()应用fn最后用原indices()重建torch.sparse_coo_tensor对 sparse CSR 则保留crow_indices/col_indices、只对values()应用fnunary.py。六、支持的二元算子Binary Operators二元算子的实现带有一个明确且保守的 caveat两个 MaskedTensor 的掩码必须匹配match否则直接抛错。这是文档强调的设计决策——为了确保用户清楚知道正在发生什么并对自己选择的掩码语义保持有意为之而不是悄悄猜测合并语义。若你确实需要某个算子的新语义请携带最小复现样例去 GitHub 提交 issue源码中的错误信息也明确指向pytorch/maskedtensor的 issue 区。当前支持的二元算子完整列表出自文档与 binary.py 的BINARY_NAMES一致add atan2 arctan2 bitwise_and bitwise_or bitwise_xor bitwise_left_shift bitwise_right_shift div divide floor_divide fmod logaddexp logaddexp2 mul multiply nextafter remainder sub subtract true_divide eq ne le ge greater greater_equal gt less_equal lt less maximum minimum fmax fmin not_equal支持inplace版本的二元算子为上述全部减去以下 6 个logaddexp logaddexp2 equal fmin minimum fmax源码中的实现要点binary.py_masks_match(*args[:2])校验掩码一致性不匹配则抛出Input masks must match. ...binary.py实际计算时把 data 解包出来对普通 Tensor 执行fn结果 mask 取自_get_at_least_one_mask——即至少一个操作数是 MaskedTensor返回其中非空的那份 maskbinary.py对 strided 布局结果 mask 会expand_as(result_data)以适配广播后的形状对稀疏布局则要求两个输入的indicesCOO或crow/col_indicesCSR一致binary.py。七、归约算子ReductionsMaskedTensor 提供以下归约算子均支持 autograd文档建议结合 Overview 与 Advanced semantics 教程理解其语义设计sum mean amin amax argmin argmax prod all norm var std这一列表与 reductions.py 中的REDUCE_NAMES完全对应且同时注册到三处NATIVE_REDUCE_MAPtorch.ops.aten级别的原生算子TORCH_REDUCE_MAPtorch.*函数如torch.sumTENSOR_REDUCE_MAPtorch.Tensor方法如t.sum()。实现机制的关键点reductions.py整体归约reduce all无dim时通过torch.masked下对应的掩码归约函数如torch.masked.sum计算结果 mask 用torch.any(mask)得到全被遮罩则为 False按维归约reduce dimresult_data masked_fn(self, dimdim, keepdimkeepdim, dtypedtype, maskself.get_mask())结果 mask 由_multidim_any(mask, dim, keepdim)沿各维做torch.any得到——输出 mask 的语义是该输出位置上至少有一个有效输入argmin/argmax 的特殊处理对于 sparse COO由于原生稀疏算子无法直接给出正确索引实现通过data.to_sparse_coo().indices()结合 strides 手工换算出一维索引reductions.py限制reduce_dim带dim的版本对稀疏布局尚未实现调用时会打印The sparse version of {fn} is not implemented in reductions.警告并返回NotImplementedreductions.py。八、视图与选择算子View and Select Functions视图/选择算子的语义最直观对 data 与 mask 分别应用算子再把结果包回 MaskedTensor。文档以select给出了完整示例 data torch.arange(12, dtypetorch.float).reshape(3, 4) data tensor([[ 0., 1., 2., 3.], [ 4., 5., 6., 7.], [ 8., 9., 10., 11.]]) mask torch.tensor([[True, False, False, True], [False, True, False, False], [True, True, True, True]]) mt masked_tensor(data, mask) data.select(0, 1) tensor([4., 5., 6., 7.]) mask.select(0, 1) tensor([False, True, False, False]) mt.select(0, 1) MaskedTensor( [ --, 5.0000, --, --] )当前支持的视图/选择算子完整列表出自文档atleast_1d broadcast_tensors broadcast_to cat chunk column_stack dsplit flatten hsplit hstack kron meshgrid narrow nn.functional.unfold ravel select split stack t transpose vsplit vstack Tensor.expand Tensor.expand_as Tensor.reshape Tensor.reshape_as Tensor.unfold Tensor.view源码层面这一族算子主要由 passthrough.py 的PASSTHROUGH_FNS承载。其模块 docstring 准确描述了设计思路这些函数应简单地同时作用于 mask 与 data——例如select、stack先对 data 应用、再对 mask 应用最后_wrap_result包回 MaskedTensorpassthrough.py。PASSTHROUGH_FNS中还包含slice、index、unsqueeze、unfold及其 backward 对应算子slice_backward、select_backward、unfold_backward、im2col/col2im等底层实现passthrough.py。九、torch.masked.maskedtensor.coreis_masked_tensor 与辅助校验文档专门为torch.masked.maskedtensor.core模块留出小节并公开了以下函数is_masked_tensoris_masked_tensor 是一个类型守卫TypeIs[MaskedTensor]返回isinstance(obj, MaskedTensor)的结果常用于实现内部判断某个参数是否是 MaskedTensor。其 docstring 给出了与工厂函数一致的用法示例 from torch.masked import MaskedTensor data torch.arange(6).reshape(2, 3) mask torch.tensor([[True, False, False], [True, True, False]]) mt MaskedTensor(data, mask) is_masked_tensor(mt) Truecore 模块还导出了_tensors_match与_masks_match两个内部校验函数前者逐元素比较两个普通 Tensor支持 sparse COO/CSR 的 indices 递归比较exactFalse时退化为torch.allclose默认rtol1e-05, atol1e-08后者比较两个 MaskedTensor 的掩码是否完全一致core.py。这两个函数既被一元/二元/归约实现复用也被测试文件直接导入验证test/test_maskedtensor.py。另外文档提到binary、creation、passthrough、reductions、unary等子模块暂以注释形式登记在文档中用于追踪它们的公开导出与实现即为本文第五至八节所述内容。十、测试与验证MaskedTensor 如何被保障正确性仓库中 test/test_maskedtensor.py共 1000 行是 MaskedTensor 的主要测试集它直接复用 PyTorch 的算子基础设施unary_ufuncs、binary_ufuncs、reduction_ops见 test/test_maskedtensor.py来生成与原生算子对齐的用例。其中有几个值得注意的验证手段_compare_mt_t把 MaskedTensor 结果与普通 Tensor 结果都做masked_fill_(~mask, 0)后再逐元素比较含rtol/atol容差从而校验遮罩后语义一致test/test_maskedtensor.py_compare_forward_backward分别对 MaskedTensor用masked_tensor(..., requires_gradTrue)与把被遮罩位置填-inf的普通 Tensor执行同一函数并反向传播同时比较 forward 结果与梯度test/test_maskedtensor.py测试覆盖了strided / sparse_coo / sparse_csr三种 layout、[] / [2] / [3,5] / [3,2,1,2]等形状以及多种 dtypetest/test_maskedtensor.py。从源码结构看这类测试构成了 MaskedTensor 语义的行为契约任何对掩码语义的改动都必须同时通过 MaskedTensor 自身与对应原生算子的双重校验。十一、原型阶段的使用须知与限制总结综合文档与源码使用 MaskedTensor 前应当了解以下几点API 处于原型阶段构造时会触发UserWarning未来可能变动core.py不建议对requires_gradTrue的 Tensor 直接构造 MaskedTensor__new__会额外告警建议先data.detach().clone()若需要可微构造应使用as_masked_tensorcore.py二元运算要求掩码严格匹配这是刻意的保守设计不匹配会直接报错算子覆盖仍在扩展中未实现的算子会在__torch_dispatch__中告警并返回NotImplemented带dim的归约暂不支持稀疏布局一元/二元算子对额外的 Tensor 位置参数有限制分别要求len(kwargs) 0、不允许第三个 Tensor 参数见 unary.py 与 binary.py布局与 dtype 有明确约束仅支持strided / sparse_coo / sparse_csr三种 layout、mask 必须为bool、data 仅支持文档所列浮点/布尔/整数 dtype。围绕上述限制如需申请新算子或新语义官方推荐携带最小可复现代码片段提交 issue并尽量附上语义提案这正是文档反复强调的让掩码语义成为一等公民协作路径的一部分。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表