如何让PyTorch Tensor始终满足自定义属性约束?
子类化torch.Tensor并确保自定义属性始终生效
问题背景
希望子类化torch.Tensor,创建一个始终满足自定义属性的张量:要求张量最后一维的和为1(表示类别概率分布)。现有实现能在实例化和__setitem__时验证,但在张量运算(如加法)或in-place操作(如fill_)时验证失效,生成不符合要求的实例。
现有实现
class ValidatedArray(torch.Tensor): def __init__(self, array: torch.Tensor): self.validate_array() def __setitem__(self, key, value): super().__setitem__(key, value) self.validate_array() def validate_array(self): assert torch.allclose(self.sum(-1), torch.ones(1)), f'The last dim represents a categorical distribution. It must sum to one.'
验证生效场景
- 实例化时验证:
>>> array = torch.ones(3,4) >>> va1 = ValidatedArray(array) AssertionError: The last dim represents a categorical distribution. It must sum to one.
- 设置元素时验证:
>>> array = torch.nn.functional.softmax(torch.ones(3,4), -1) >>> va1 = ValidatedArray(array) >>> va1[0] = 1 AssertionError: The last dim represents a categorical distribution. It must sum to one.
验证失效场景
- 张量运算后生成非法实例:
>>> array = torch.nn.functional.softmax(torch.ones(3,4), -1) >>> va1 = ValidatedArray(array) >>> va2 = va1 + 2 >>> va2 ValidatedArray([[2.2500, 2.2500, 2.2500, 2.2500], [2.2500, 2.2500, 2.2500, 2.2500], [2.2500, 2.2500, 2.2500, 2.2500]])
- in-place操作后生成非法实例:
>>> array = torch.nn.functional.softmax(torch.ones(3,4), -1) >>> va1 = ValidatedArray(array) >>> va1.fill_(2.) >>> va1 ValidatedArray([[2., 2., 2., 2.], [2., 2., 2., 2.], [2., 2., 2., 2.]])
解决方案
要确保所有操作后ValidatedArray实例都满足验证条件,需要利用PyTorch张量子类的特性,从以下几个方面改进:
1. 重写__new__方法完成实例化验证
PyTorch张量的实例创建依赖__new__方法,而非__init__,因此需要在__new__中完成初始验证:
class ValidatedArray(torch.Tensor): @classmethod def __new__(cls, array: torch.Tensor): # 创建张量实例 instance = super().__new__(cls, array) # 执行验证 instance.validate_array() return instance def __init__(self, array: torch.Tensor): # __init__在实例创建后调用,此处可留空或做其他初始化 pass def __setitem__(self, key, value): super().__setitem__(key, value) self.validate_array() def validate_array(self): # 放宽浮点容差避免精度问题 assert torch.allclose(self.sum(-1), torch.ones_like(self.sum(-1)), atol=1e-6), \ f'The last dim represents a categorical distribution. It must sum to one.'
2. 拦截所有Torch操作,验证结果
通过重写__torch_function__方法,拦截所有针对ValidatedArray的Torch操作,在返回结果前验证是否符合要求:
@classmethod def __torch_function__(cls, func, types, args=(), kwargs=None): if kwargs is None: kwargs = {} # 执行原操作 result = super().__torch_function__(func, types, args, kwargs) # 如果结果是ValidatedArray实例,执行验证 if isinstance(result, ValidatedArray): result.validate_array() return result
3. 处理in-place操作
部分in-place操作(如fill_)不会触发__torch_function__拦截,需要单独重写这些方法,在操作后执行验证:
def fill_(self, value): super().fill_(value) self.validate_array() return self # 按需添加其他in-place方法的重写,如add_、mul_等 def add_(self, other): super().add_(other) self.validate_array() return self
完整实现
import torch class ValidatedArray(torch.Tensor): @classmethod def __new__(cls, array: torch.Tensor): instance = super().__new__(cls, array) instance.validate_array() return instance def __init__(self, array: torch.Tensor): pass def __setitem__(self, key, value): super().__setitem__(key, value) self.validate_array() def validate_array(self): assert torch.allclose(self.sum(-1), torch.ones_like(self.sum(-1)), atol=1e-6), \ f'The last dim represents a categorical distribution. It must sum to one.' @classmethod def __torch_function__(cls, func, types, args=(), kwargs=None): if kwargs is None: kwargs = {} result = super().__torch_function__(func, types, args, kwargs) if isinstance(result, ValidatedArray): result.validate_array() return result def fill_(self, value): super().fill_(value) self.validate_array() return self def add_(self, other): super().add_(other) self.validate_array() return self
验证改进后的效果
- 张量运算会触发验证:
>>> array = torch.nn.functional.softmax(torch.ones(3,4), -1) >>> va1 = ValidatedArray(array) >>> va2 = va1 + 2 AssertionError: The last dim represents a categorical distribution. It must sum to one.
- in-place操作会触发验证:
>>> array = torch.nn.functional.softmax(torch.ones(3,4), -1) >>> va1 = ValidatedArray(array) >>> va1.fill_(2.) AssertionError: The last dim represents a categorical distribution. It must sum to one.
关键说明
- PyTorch张量子类的实例化逻辑依赖
__new__,必须在此阶段完成初始验证,否则无法覆盖所有实例创建场景。 __torch_function__是PyTorch提供的统一拦截机制,能覆盖绝大多数张量操作,但部分in-place方法需要单独重写。- 验证时建议设置合理的浮点容差(如
atol=1e-6),避免因浮点精度问题误判。
内容的提问来源于stack exchange,提问作者Doug Tischer
相关产品推荐
相关产品推荐

