修改detach得到的numpy数组为何会改变PyTorch模型的原始权重?
问题原因解析
这个现象是PyTorch张量和NumPy数组的内存共享机制导致的,具体逻辑如下:
- PyTorch中所有存储在CPU上的张量,调用
.numpy()方法得到的NumPy数组,默认和原张量共享同一个底层内存缓冲区,对其中任意一个做原地修改,都会同步修改另一个的值。 - 逐行拆解你的代码执行逻辑:
x[0][1]是你模型的conv1.weight参数,属于PyTorch的Parameter类(Tensor的子类),存储着模型的原始权重。.detach()方法仅将张量从计算图中剥离,取消梯度追踪关联,不会拷贝底层数据,返回的张量仍然和原始参数共享内存。.cpu()方法如果原始张量在GPU上,会将数据拷贝到CPU生成新的CPU张量;如果原始张量本身就在CPU上,返回的张量仍然和原始参数共享内存。.numpy()将CPU张量转换为NumPy数组,这一步默认共享内存,不会生成数据副本。.squeeze()仅对数组的维度做压缩,返回的是原数组的视图,没有数据拷贝操作,最终得到的weight数组仍然和模型的原始权重共享底层内存。
- 你执行
weight[0] = weight[0] * 0 + 0.1属于原地修改操作,直接修改了共享内存中的数据,因此模型的原始参数会同步发生变化。
如果需要修改NumPy数组时不影响原始模型权重,只需要在转换时显式拷贝数据即可,代码修改为:
weight = x[0][1].detach().cpu().numpy().squeeze().copy()
调用.copy()后会生成独立的NumPy数组,和原张量不再共享内存,修改时不会影响模型原始参数。
内容的提问来源于stack exchange,提问作者Sarvagya Gupta
相关产品推荐
相关产品推荐

