含numpy数组属性的类对象复制同步修改问题
解决Python类中numpy数组浅拷贝导致的同步修改问题
问题复现
你定义的Test类及测试代码如下:
import numpy as np import copy as cp class Test(): x=5 arr=np.array([5,3]) TEST_1=Test() TEST_2=cp.copy(TEST_1)
修改TEST_1的x属性时,TEST_2的x不受影响:
print(TEST_2.x) TEST_1.x+=1 print(TEST_2.x)
输出:
>>>5 >>>5
但修改TEST_1的arr数组时,TEST_2的arr会同步改变:
print(TEST_2.arr) TEST_1.arr+=np.array([1,1]) print(TEST_2.arr)
输出:
>>> [5 3] >>> [6 4]
问题原因
- 类属性的共享特性:当前
x和arr是Test类的类属性,所有实例默认共享。修改TEST_1.x时,Python会自动为TEST_1创建同名实例属性,覆盖类属性,所以TEST_2仍引用类属性的x,不受影响;但修改TEST_1.arr时,并没有创建新的实例属性,而是直接修改了类属性指向的numpy数组(可变对象),所有实例都会同步看到变化。 - 浅拷贝的局限性:
copy.copy()是浅拷贝,仅复制对象引用,对于numpy数组这类可变对象,不会复制数组本身,两个实例仍共享同一个数组对象。
解决方案
方案1:改为实例属性+深拷贝
重写Test类,在__init__中初始化实例属性,同时用copy.deepcopy()实现完全拷贝:
import numpy as np import copy as cp class Test(): def __init__(self): # 每个实例拥有独立的实例属性 self.x = 5 self.arr = np.array([5,3]) TEST_1 = Test() # 深拷贝确保数组也被完整复制 TEST_2 = cp.deepcopy(TEST_1)
测试验证:
print(TEST_2.arr) TEST_1.arr += np.array([1,1]) print(TEST_2.arr)
输出:
>>> [5 3] >>> [5 3]
方案2:自定义类的__copy__方法
如果需要用浅拷贝但希望数组独立,可以重写__copy__方法,手动复制数组:
import numpy as np import copy as cp class Test(): def __init__(self): self.x = 5 self.arr = np.array([5,3]) def __copy__(self): new_instance = Test() new_instance.x = self.x # 用numpy数组自身的.copy()方法创建副本 new_instance.arr = self.arr.copy() return new_instance TEST_1 = Test() TEST_2 = cp.copy(TEST_1)
此时用copy.copy()也能得到独立的数组实例,修改其中一个不会影响另一个。
方案3:手动复制数组(不修改类)
如果不想改动类结构,可在复制实例后手动替换数组:
TEST_1 = Test() TEST_2 = cp.copy(TEST_1) # 手动复制数组,让TEST_2拥有独立的数组对象 TEST_2.arr = TEST_1.arr.copy()
内容的提问来源于stack exchange,提问作者Lory1502
相关产品推荐
相关产品推荐

