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

为什么PyTorch克隆张量会移除自定义属性,对应的解决方法是什么?

问题原因

PyTorch 中张量的 clone() 是C++后端实现的接口,默认仅复制张量核心属性:包括形状、数据类型、存储设备、requires_grad标记等官方预置的元数据,用户在Python层手动绑定的自定义属性属于Python对象的动态附加属性,没有被纳入克隆逻辑的处理范围,所以不会被复制到新张量上。
另外你示例中待克隆的a是nn.Parameter类型,调用clone()后默认返回的是普通torch.Tensor对象,类型本身发生了变化,也会导致原Parameter上绑定的自定义属性全部丢失。

解决方法

你可以根据自己的使用场景选择以下方案:

  • 方案1:手动复制属性,适合自定义属性数量少的场景
import torch
from torch import nn

a = nn.Parameter(torch.rand(1))
a.adapt = True                      # 定义自定义属性

b = a.clone()
b = nn.Parameter(b)                 # 保持Parameter类型(按需选择)
b.adapt = a.adapt                   # 手动绑定自定义属性

print(b.adapt) # 输出 True
  • 方案2:封装通用带属性克隆函数,适合有多个自定义属性的场景
import torch
from torch import nn

def clone_with_custom_attr(tensor):
    # 先克隆张量本体
    cloned_tensor = tensor.clone()
    # 如果原对象是Parameter,克隆后也保持Parameter类型
    if isinstance(tensor, nn.Parameter):
        cloned_tensor = nn.Parameter(cloned_tensor)
    # 遍历复制所有用户自定义属性
    default_attrs = set(dir(type(tensor)))
    for attr_name, attr_val in tensor.__dict__.items():
        if attr_name not in default_attrs:
            setattr(cloned_tensor, attr_name, attr_val)
    return cloned_tensor

# 测试
a = nn.Parameter(torch.rand(1))
a.adapt = True
a.custom_num = 1024

b = clone_with_custom_attr(a)
print(b.adapt) # 输出 True
print(b.custom_num) # 输出 1024
  • 方案3:自定义张量子类,适合需要长期大量使用带自定义属性张量的场景
import torch

class AttrTensor(torch.Tensor):
    @staticmethod
    def __new__(cls, data, requires_grad=False, **custom_attrs):
        instance = super().__new__(cls, data, requires_grad=requires_grad)
        # 初始化自定义属性
        for k, v in custom_attrs.items():
            setattr(instance, k, v)
        return instance
    
    def clone(self, *args, **kwargs):
        cloned = super().clone(*args, **kwargs)
        # 克隆时自动复制所有自定义属性
        default_attrs = set(dir(torch.Tensor))
        for attr_name, attr_val in self.__dict__.items():
            if attr_name not in default_attrs:
                setattr(cloned, attr_name, attr_val)
        return cloned

# 测试
a = AttrTensor(torch.rand(1), requires_grad=True, adapt=True)
b = a.clone()
print(b.adapt) # 输出 True

内容的提问来源于stack exchange,提问作者Blade

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 21:45:05