Python函数中numpy数组参数传递问题:np.dot结果无法返回的解决方法
numpy数组参数传递:np.dot无法修改外部数组的解决方法
我在用Python实现矩阵运算,用numpy数组存储矩阵。已知标量参数按值传递,数组参数类似C的指针传递。示例1手动实现矩阵乘法能把结果传回调用端,但示例2用np.dot替代后,外部数组没被修改,怎么解决?
示例1(手动乘法,有效)
import numpy as np def my_func(a,b,c): for i in range (0,2): for j in range (0,2): c[i,j]=0. for k in range(0,2): c[i,j]=c[i,j]+a[i,k]*b[k,j] print("c") print(c) d = np.array([[1,2],[3,4]]) e = np.array([[-1,3],[2,-1]]) f = np.zeros((2,2)) my_func(d,e,f) print("f") print(f)
输出:
c [[3. 1.] [5. 5.]] f [[3. 1.] [5. 5.]]
示例2(np.dot,无效)
import numpy as np def my_func(a,b,c): c=np.dot(a,b) print("c") print(c) d = np.array([[1,2],[3,4]]) e = np.array([[-1,3],[2,-1]]) f = np.zeros((2,2)) my_func(d,e,f) print("f") print(f)
输出:
c [[3 1] [5 5]] f [[0. 0.] [0. 0.]]
问题原因
示例1是原地修改数组的元素:函数里的c[i,j] = ...直接操作传入数组的内存空间,所以外部的f能同步看到变化。
但示例2里的c = np.dot(a,b)是给函数内部的局部变量c重新赋值了一个全新的numpy数组——原来指向外部f的引用被覆盖了,后续操作和外部的f完全无关,所以外部数组的值没变化。
解决方法
方法1:原地写入结果(修改原数组内容)
用切片赋值c[:] = np.dot(a,b),把np.dot的结果写入原数组的内存空间,而不是重新赋值变量:
import numpy as np def my_func(a,b,c): c[:] = np.dot(a,b) print("c") print(c) d = np.array([[1,2],[3,4]]) e = np.array([[-1,3],[2,-1]]) f = np.zeros((2,2)) my_func(d,e,f) print("f") print(f)
输出:
c [[3 1] [5 5]] f [[3. 1.] [5. 5.]]
方法2:函数返回结果,外部接收
直接让函数返回np.dot的结果,外部变量接收后覆盖原数组:
import numpy as np def my_func(a,b): result = np.dot(a,b) print("result") print(result) return result d = np.array([[1,2],[3,4]]) e = np.array([[-1,3],[2,-1]]) f = np.zeros((2,2)) f = my_func(d,e) print("f") print(f)
输出:
result [[3 1] [5 5]] f [[3 1] [5 5]]
方法3:用np.dot的out参数
np.dot支持out参数,直接将计算结果写入指定的输出数组:
import numpy as np def my_func(a,b,c): np.dot(a,b, out=c) print("c") print(c) d = np.array([[1,2],[3,4]]) e = np.array([[-1,3],[2,-1]]) f = np.zeros((2,2)) my_func(d,e,f) print("f") print(f)
输出:
c [[3. 1.] [5. 5.]] f [[3. 1.] [5. 5.]]
内容的提问来源于stack exchange,提问作者Mike Gunn
相关产品推荐
相关产品推荐

