在NumPy中实现代价函数时,如何区分使用np.dot与*运算符?
NumPy中np.dot与*运算符的区别及使用判断
为什么公式中用np.dot(w.T, X)而非w.T * X
首先明确公式里的w.TX是线性代数中的矩阵乘法(向量点积):
- 假设变量维度符合机器学习常规定义:
w是(n, 1)的权重列向量,X是(n, m)的特征矩阵(每列对应1个样本,共m个样本,每个样本含n个特征)。 w.T将w转置为(1, n)的行向量,此时w.T与X做矩阵乘法,满足前一矩阵列数(n)等于后一矩阵行数(n)的规则,结果是(1, m)的行向量,每个元素对应单个样本的线性组合w.Tx(i),完全匹配公式1的计算逻辑。
而*在NumPy中是逐元素乘法,要求两个数组形状完全匹配(或符合广播规则且逐元素对应):
- 对于
(1, n)的w.T和(n, m)的X,两者形状不兼容,直接用*会抛出ValueError;即便形状巧合兼容,得到的也是逐元素相乘的结果,完全不是公式需要的线性组合,不符合计算要求。
如何判断何时用np.dot,何时用*
- 用
np.dot()(或Python3.5+支持的@运算符,效果等价)的场景:- 需要计算线性代数中的矩阵乘法,比如权重与特征的线性组合、矩阵间的乘法运算
- 需要计算向量点积,比如两个同维度向量的内积
- 用
*运算符(或np.multiply())的场景:- 需要对两个数组进行逐元素对应相乘,比如代价函数中
Y*np.log(A),这里Y和A都是(m,)的向量,逐元素相乘后得到每个样本的y(i)log(a(i)),再求和就符合公式2的累加逻辑
- 需要对两个数组进行逐元素对应相乘,比如代价函数中
内容的提问来源于stack exchange,提问作者kayalotta
相关产品推荐
相关产品推荐

