为什么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
相关产品推荐
相关产品推荐

