使用NumPy实现ReLU导数时,d_relu为何会修改y变量?
为什么d_relu函数会修改外部的numpy数组y?
问题描述
我编写了一段测试ReLU激活函数及其导数的代码:
import numpy as np def relu(z): return np.maximum(0,z) def d_relu(z): z[z>0]=1 z[z<=0]=0 return z x=np.array([5,1,-4,0]) y=relu(x) z=d_relu(y) print("y = {}".format(y)) print("z = {}".format(z))
运行后输出为 y = [1 1 0 0] z = [1 1 0 0],但我预期的结果应该是 y = [5 1 0 0] z = [1 1 0 0]。我原本以为函数调用是值传递(传递变量的副本),为什么d_relu函数会对外部的y变量产生修改呢?
问题解析
首先要纠正一个常见误区:Python的参数传递既不是纯粹的值传递,也不是纯粹的引用传递,而是「传对象引用」。针对不同类型的对象,表现会有所不同:
- 对于不可变对象(比如int、字符串、tuple),在函数内修改会创建新的对象,不会影响外部的原变量;
- 对于可变对象(比如numpy数组、list、dict),如果在函数内直接修改对象本身的内容(而不是给参数重新赋值一个新对象),那么外部的原对象也会被改变——因为函数参数和外部变量指向的是同一个底层对象。
回到你的代码一步步看:
y = relu(x):np.maximum(0,z)会返回一个新的numpy数组,此时y指向这个数组[5,1,0,0];- 调用
d_relu(y)时,我们把y指向的数组的引用传递给了参数z——也就是说,函数里的z和外部的y,指向的是同一个numpy数组对象; - 在
d_relu中,z[z>0]=1和z[z<=0]=0这两行是直接修改数组对象的元素值,而不是给z重新赋值一个新数组; - 所以当函数执行完毕后,外部的y自然也会显示被修改后的数组内容。
解决方案
要避免修改原数组,我们可以在d_relu函数里先创建输入数组的副本,再对副本进行修改,这样就不会影响原对象了:
修改后的d_relu函数:
def d_relu(z): # 创建输入数组的副本,避免修改原数组 z_copy = z.copy() z_copy[z_copy>0] = 1 z_copy[z_copy<=0] = 0 return z_copy
现在再运行代码,就能得到你预期的输出:y = [5 1 0 0] z = [1 1 0 0]
另外补充一点:如果你在函数里是给z重新赋值(比如z = np.where(z>0, 1, 0)),那也不会影响外部的y,因为这时候是让函数里的z指向了一个新的对象,原对象并没有被改变。而你原来的写法是直接修改数组的元素,所以才会影响外部变量。
内容的提问来源于stack exchange,提问作者Egbert
相关产品推荐
相关产品推荐

