MNIST实践中W与x点积的argmax含义及向量追加1的原因
MNIST深度学习入门代码问题解答
关于预测代码中argmax的实际作用
首先明确:这个操作不是取权重矩阵的最大值,你的初步理解这里存在偏差。
这段代码是MNIST 10分类任务(识别0-9共10个手写数字)的单层线性分类预测逻辑,逐行拆解逻辑如下:
- 权重矩阵
W的形状为(10, 785),10行分别对应10个数字类别的权重参数,785列对应输入特征的维度。 - 单样本
x是长度785的一维特征向量,np.dot(W, x)的计算结果是长度为10的一维数组,数组中第i位的数值,就是模型计算出的「该样本属于第i类」的置信度得分。 np.argmax()的作用是找出这个长度为10的得分数组中,数值最大的元素对应的索引,这个索引就是模型输出的预测类别。
举个直观例子:如果某张图计算得到的点积结果是[0.02, 0.01, 0.03, 0.9, 0.01, 0.01, 0.005, 0.01, 0.005, 0.0],np.argmax会返回索引3,代表模型预测这张图是数字3。
关于图像展平后追加1的设计原因
这个追加的固定值1是偏置项的统一计算技巧,是线性类模型的常规写法:
- 线性分类器的原始计算式是
score = W·x + b,其中b是偏置向量,作用是给每个类别的得分加一个固定偏移,让分类决策边界不需要强制经过坐标原点,降低模型拟合的约束。 - 如果严格按照原始公式写代码,前向传播需要先算权重和特征的点积,再单独加一次偏置,参数更新时也要单独维护偏置的更新逻辑,比较繁琐。
- 给所有输入特征末尾追加一个固定为1的维度后,我们只需要把原本独立的偏置参数拼接到权重矩阵
W的最后一列,就可以把原本的「点积+加偏置」两步计算,合并成一次np.dot(W_new, x_new)的点积运算:点积计算时,新增的1会和W最后一列的偏置参数相乘,刚好等价于原始公式里单独加偏置的步骤,计算结果完全一致,但代码逻辑简化了很多。
你看到的代码里只给输入追加了1,说明权重初始化时已经把偏置对应的参数位包含进去了(W的列数是785,刚好对应784个像素特征+1个偏置位),不需要额外单独定义、更新偏置参数。
内容的提问来源于stack exchange,提问作者pav
相关产品推荐
相关产品推荐

