NumPy与PyTorch跨步操作:如何实现内存共享的原地写入?
问题解答
1. PyTorch实现目标累加操作的方法
PyTorch报错是因为原地操作(+=)中输入张量s与目标张量t共享内存,PyTorch严格禁止这种可能导致计算歧义的内存重叠操作。要实现你想要的第一行自增后,第二行与修改后的第一行累加的效果,分两步原地操作即可:
import torch t = torch.arange(6.).reshape(2,3) + 3 # 第一步:第一行原地自增(等价于 t[0] = t[0] * 2) t[0] += t[0] # 第二步:第二行与修改后的第一行累加 t[1] += t[0] print(t) # 输出: # tensor([[ 6., 8., 10.], # [12., 15., 18.]])
如果一定要保留as_strided的用法,需要先确保操作时不共享内存(比如先克隆s的值),但要得到目标结果依然需要分两步:
import torch t = torch.arange(6.).reshape(2,3) + 3 # 获取第一行的视图 s = t.as_strided(size=(1,3), stride=(t.stride(0), t.stride(1))) # 第一步:第一行自增 t[0] += s.clone() # 第二步:第二行累加修改后的第一行 t[1] += t[0]
2. NumPy不复制数据实现目标累加的方法
你的原始NumPy代码执行array += s时,NumPy会先将s广播为与array同形状的临时数组(基于s的原始值),再执行原地加法,因此第二行累加的是第一行的原始值,而非修改后的值。要实现目标效果,同样需要分两步原地操作,全程无需复制数据:
import numpy as np array = np.arange(6.).reshape(2,3) + 3 # 第一步:第一行原地自增 array[0] += array[0] # 第二步:第二行与修改后的第一行累加 array[1] += array[0] print(array) # 输出: # [[ 6. 8. 10.] # [12. 15. 18.]]
这里的array[0]和array[1]都是原数组的视图,所有操作都是直接修改原始内存,没有数据复制。
内容的提问来源于stack exchange,提问作者Yorai Levi
相关产品推荐
相关产品推荐

