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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:50:25