1. 鱼书里的“算术计算”到底是什么
第一次看到“鱼书-python算术计算”这个标题,我猜你十有八九也是冲着《深度学习入门:基于Python的理论与实现》这本书来的。“鱼书”就是圈内对斋藤康毅那本封面上有鱼的书的爱称。这本书我从头到尾撸了两遍,第一遍当速读,第二遍才真正把每个算术细节抠明白。
这本书最大的特点,是不用现成的深度学习框架,而是用Python和NumPy从零实现神经网络。也就是说,你写的每一行代码,都得亲手完成加法、乘法、矩阵运算、指数、对数、除法这些最基础的算术计算。很多人读这本书卡壳,不是卡在神经网络原理上,而是卡在“算术计算”上。
为什么这么说?因为书里涉及的算术计算,远不是小学口算那种程度,而是带着维度、带着广播、带着数据类型的批量运算。比如:
- 数组和标量相乘,结果是数组逐元素乘标量;
- 两个二维数组做乘法,区分“对应元素相乘”和“矩阵乘法”;
- 求损失函数时,要对一批样本的误差求平均值;
- 反向传播时,要对中间变量做微积分意义上的梯度计算,实际落到代码里就是四则运算的组合。
换句话说,鱼书的算术计算,是Python入门到NumPy进阶再到深度学习原理的一条链条。你要是不把这块啃下来,后面Softmax、交叉熵、梯度下降全都像看天书。
这篇文章适合三类人:刚装好Python还没玩熟NumPy的新手、在鱼书某些公式和代码之间来回横跳的读者、以及想系统补一补Python算术计算底子的同学。我会把我在实操里踩过的坑、验证过的方法、以及为什么这样算的道理都揉碎了讲。
2. 先理顺Python原生的算术运算基本功
2.1 变量、类型与运算符的底层逻辑
Python的算术计算,第一步不是敲公式,而是搞清楚“参与计算的东西到底是什么”。很多人写代码报错,十次里有八次是类型问题。鱼书里的代码大量涉及标量、一维数组、二维数组、布尔值之间的混合运算,稍不注意类型,结果就完全不是你想的那样。
Python原生支持的基本算术运算符有这些:
+:加法,两个数值相加,字符串序列也能拼接-:减法*:乘法,数组和列表用乘法意义完全不同/:除法,结果为浮点数//:整除,向下取整%:取余**:幂运算
一个非常容易翻车的地方是/和//的区别。比如7 / 2在Python里得到3.5,而7 // 2得到3。鱼书里计算准确率、损失值这些指标时,如果你用了整除,小数点全被砍掉,训练曲线看起来就是一根根锯齿,根本没法判断模型有没有收敛。我初学时就干过这种事,拿//去算平均损失,结果100个epoch的损失全是0和1在跳,排查了半天才发现是整除的问题。
整数和浮点数之间的转换也是大量算术计算里的隐性坑。Python里整数和浮点数运算,结果会自动提升为浮点数,这个特性叫“数值提升”。比如1 + 2.0结果是3.0而不是3。但在某些场景下,比如用整数数组索引、或者把结果作为数组的shape参数时,浮点数会出现意外报错。鱼书里常见的场景是int(round(...))的写法,就是为了避免浮点数直接进索引或shape。
2.2 常用内置算术函数及其等价写法
除了运算符,Python还内置了一批数学函数,主要集中在math模块和几个内置函数上:
abs(x):求绝对值。鱼书里算误差、算梯度的范数时经常用round(x, n):四舍五入到n位小数pow(x, y):等价于x ** ymax和min:返回最大值和最小值sum:对列表或可迭代对象求和
这些函数看起来基础,但在鱼书配套的例子里,它们的组合能实现很多功能。比如你要判断一个神经网络的预测是否正确,最简单的做法就是np.argmax(y)找到最大值的索引,再和标签对比。argmax虽然来自NumPy,但它的思想就是对一组数做比较计算,这和内置的max是同一类逻辑。
我建议你把math模块里的这些函数也过一遍:math.exp、math.log、math.sqrt、math.floor、math.ceil。鱼书在讲激活函数和损失函数时,数学公式用的就是指数和对数,虽然最终代码层面会用NumPy实现,但理解原生的实现方式有助于你搞清楚每一步数学含义。
2.3 一个容易忽略的细节:运算优先级
算术计算还有一个基础中的基础——运算顺序。Python遵守常规数学优先级:**先于* /先于+ -,括号可以改变顺序。但有一种写法是新手特别容易看懵的,就是复合赋值运算,比如x += 1、y *= 2。在鱼书的代码里,参数的更新经常写成:
W -= learning_rate * grad这实际上是W = W - learning_rate * grad的简写。逻辑很清楚,但如果你之前没有养成拆解的习惯,可能看不出来这里发生了“先乘后减”的顺序。等你自己写优化器更新参数时,只要漏看一个符号,整个训练就废了。
我的建议是:在所有涉及算术计算的代码里,即使你确定优先级正确,也尽量用括号把意图写清楚。代码不仅是给机器跑的,更是给人看的,尤其你在读鱼书这种代码密度很高的出版物时,括号能帮你省掉大量脑力。
3. NumPy算术计算才是主角
3.1 为什么鱼书要引入NumPy而不是纯Python列表
纯Python的列表当然也可以做算术计算,比如[1, 2, 3]和[4, 5, 6]想逐元素相加,可以用列表推导式:
[x + y for x, y in zip([1, 2, 3], [4, 5, 6])]结果得到[5, 7, 9]。但是,这种写法有两大问题:一是代码冗长,二是一旦维度升到二维、三维,列表推导式嵌套起来就没法看了。鱼书里的图像数据动辄是(60000, 784)这种形状的数组,你要是用列表推导去算批量数据的均值或梯度,性能会差到怀疑人生。
NumPy的核心是ndarray——一个高性能的多维数组对象。它比Python列表强的点在于:
- 运算由C语言底层实现,循环速度远超Python逐元素遍历
- 支持广播机制,不同形状的数组也能算术运算
- 自带全套数学函数,如
np.exp、np.log、np.sum、np.mean - 支持切片、索引、重塑等操作,数据预处理极其方便
安装NumPy的方式,我实测最稳的是用pip:
pip install numpy如果网络环境一般,可以指定国内镜像源加速:
pip install numpy -i https://pypi.tuna.tsinghua.edu.cn/simple3.2 数组的基本算术运算:逐元素与矩阵乘法的区别
NumPy里最需要分清的算术计算有两种:逐元素乘法和矩阵乘法。前者用*,后者用np.dot()或@运算符。看个具体例子:
import numpy as np a = np.array([[1, 2], [3, 4]]) b = np.array([[5, 6], [7, 8]]) # 逐元素乘法:对应位置相乘 c = a * b # 结果 [[5, 12], [21, 32]] # 矩阵乘法:行与列做内积 d = np.dot(a, b) # 或 d = a @ b # 结果 [[19, 22], [43, 50]]这个区别在鱼书里太关键了。神经网络的全连接层计算Y = W @ X + b,用的是矩阵乘法;而激活函数的计算Y = sigmoid(Z),用的是逐元素运算。如果你把这两个搞混,网络的每一层输出都是错的,且错误会层层累积,最终训练出来的模型完全不能用。
我的判断标准很简单:凡是公式里出现求和符号或者内积概念的,基本是矩阵乘法;凡是“对每个元素单独做变换”的,就是逐元素运算。
3.3 广播机制:不同形状数组做算术计算的规则
广播机制是NumPy算术计算里最强大、也最容易踩坑的特性。它的核心规则是:当两个数组形状不同时,NumPy会自动把较小的数组扩展到较大数组的形状,再逐元素运算。
广播的原则有三条:
- 如果两个数组维度不同,把维度较少的数组形状前面补1
- 比较两个数组每个维度上的大小,要么相等,要么其中一个为1,否则无法广播
- 如果某个维度为1,则沿该维度复制
举个例子:
a = np.array([[1, 2, 3], [4, 5, 6]]) # 形状 (2, 3) b = np.array([10, 20, 30]) # 形状 (3,) c = a + b # 广播后 b 变为 [[10,20,30],[10,20,30]] # 结果 [[11, 22, 33], [14, 25, 36]]鱼书里大量代码都用到了广播。最典型的就是偏置b加在矩阵乘法结果上:W @ X + b,前者形状是(batch_size, output_size),后者形状是(output_size,),NumPy会自动把偏置加到每一行上。这个设计大大简化了代码,不用你再写循环。
但是广播机制也有坑。当维度较少数组前补1后仍然不匹配时,就会报错“operands could not be broadcast together with shapes ...”。我举一个真实案例:我在算两层神经网络中间层的梯度时,dx的形状是(3, 4),想加上一个形状为(3, 1)的向量,结果报广播错误。原因是我忘了中间变量经过转置后维度换了位置,导致两个数组在某个维度上既不是1也不相等。排查这类问题最简单的办法是把涉及的数组shape逐行打印出来,看哪一维匹配不上。
3.4 常用NumPy算术函数速查
鱼书里高频出现的NumPy算术函数,我把它们整理成一张速查表,方便你写代码时对照:
| 函数 | 功能 | 鱼书典型场景 |
|---|---|---|
np.sum(axis=...) | 求和,axis指定求和方向 | 求损失函数汇总 |
np.mean(axis=...) | 求平均值 | 批量数据的平均损失 |
np.exp(x) | 逐元素自然指数 | Softmax函数前向传播 |
np.log(x) | 逐元素自然对数 | 交叉熵损失计算 |
np.dot(a, b) | 矩阵乘法 | 全连接层输出计算 |
np.max(axis=...) | 最大值 | Softmax数值稳定性处理 |
np.argmax(axis=...) | 最大值的索引 | 分类准确率判断 |
np.sqrt(x) | 开平方 | 标准差、归一化计算 |
np.clip(x, min, max) | 数值裁剪 | 防止梯度爆炸 |
以np.sum为例,它的axis参数你一定要亲自试一遍:
a = np.array([[1, 2, 3], [4, 5, 6]]) print(a.sum()) # 所有元素和 = 21 print(a.sum(axis=0)) # 沿行方向压缩,结果 [5, 7, 9] print(a.sum(axis=1)) # 沿列方向压缩,结果 [6, 15]axis的直观理解是“我要沿着哪个轴把数组压缩掉”。鱼书里计算一批数据的交叉熵损失时,先对每个样本的各个类别做axis=1的求和,再对所有样本做axis=0的均值,逻辑非常清晰。
3.5 数组形状改变对算术计算的影响
算术计算中,数组的reshape、transpose、flatten这些操作往往直接决定后续计算能不能成立。鱼书里最让我印象深刻的一个场景是手写数字识别数据集的预处理:原始图像是28×28像素,读入后需要拉平成784维的向量,才能作为网络输入。
X = X.reshape(-1, 784)这里的-1表示自动推断该维度大小。如果你有60000张图片,这行代码就把形状变成(60000, 784)。这个操作不做,后面的矩阵乘法直接报维度不匹配错误。
transpose同样需要重视。反向传播中经常要对权重矩阵做转置,比如梯度从输出层传到隐藏层时,公式是dx = W.T @ dout。.T这个属性就是转置。如果你漏了转置,形状不匹配是小事,更可怕的是形状碰巧匹配但计算完全错误——这种情况我不止一次遇到过,当时真觉得比报错还难排查。后来我养成了一个习惯:每次做矩阵运算前,先在注释里写好关键变量的预期形状,宁可多写几行注释,也不要让形状跟丢了。
4. 鱼书中经典算术计算的实战拆解
4.1 均方误差与交叉熵的算术实现
鱼书里讲了两种损失函数:均方误差和交叉熵误差。这两种函数看似复杂,拆到底层就是一套算术流程。
均方误差的公式是:
def mean_squared_error(y, t): return 0.5 * np.sum((y - t) ** 2)这里y是神经网络的输出,t是监督标签。代码里的算术链条是:先求差,再逐元素平方,再求和,最后乘0.5。为什么乘0.5?纯粹是为了后面求导时把2次方的系数约掉,让导数表达式更干净。这属于数学上的约定,不是必须,但鱼书沿用这个习惯,你也就跟着写。
交叉熵误差公式是:
def cross_entropy_error(y, t): delta = 1e-7 return -np.sum(t * np.log(y + delta))delta是防止log(0)导致负无穷的保护项。这里的算术逻辑是:np.log(y + delta)逐元素取对数,再和标签t逐元素相乘,最后求和并取负。我每次看这个函数都觉得,它把“信息量”这种抽象概念,转化成了非常具体的数值运算。
4.2 数值微分:用算术计算逼近导数
鱼书在介绍反向传播之前,先用数值微分验证了梯度计算的正确性。所谓数值微分,就是用极限的近似形式——中心差分:
def numerical_gradient(f, x): h = 1e-4 grad = np.zeros_like(x) for idx in range(x.size): tmp_val = x[idx] x[idx] = tmp_val + h fxh1 = f(x) x[idx] = tmp_val - h fxh2 = f(x) grad[idx] = (fxh1 - fxh2) / (2 * h) x[idx] = tmp_val return grad这段代码的算术核心是:(f(x+h)-f(x-h))/(2h)。为什么用中心差分而不是前向差分(f(x+h)-f(x))/h?因为中心差分的误差是O(h²),前向差分是O(h),前者精度高一个量级。h取1e-4也是经验值——太大则截断误差不明显,太小则因浮点数精度问题产生舍入误差。我试过h=1e-8,结果梯度反而抖动,就是因为舍入误差开始主导了。
4.3 梯度下降的算术更新
梯度下降的参数更新是鱼书里算术最紧凑的地方。每次都写:
W -= learning_rate * grad拆开看就是四步:学习率乘梯度、得到增量、从当前权重减去增量、结果写回权重。整个过程没有复杂的数学,但它的意义在于把“沿着梯度反方向走一小步”这个直觉,变成了可执行的算术运算。
学习率怎么选?这是实战里最常见的超参数。鱼书给的标准值是0.1或0.01。我自己的实测经验是:学习率太大,损失会震荡甚至发散;太小,收敛慢到让人失去耐心。解决办法是先用一个小实验——同样的网络结构,分别用0.1、0.01、0.001跑50个epoch,看损失曲线的下降趋势,选那个又稳又快的。
梯度下降里还有一个细节:批量数据的梯度需要求平均。如果一次读入100张图片,每个样本都算出梯度,最终更新的方向应该是所有梯度的平均值。这个平均操作就是一个sum / batch_size的算术计算。如果你忘了除以批量大小,梯度方向虽然不变,但步长会成倍放大,模型大概率不收敛。
4.4 Softmax的算术技巧
鱼书的Softmax函数是一个绝佳的算术优化案例。看原始的数学公式,它要对每个输入求指数,然后除以所有指数之和。但直接这么算有一个致命问题:当输入里有较大数值,比如1000,np.exp(1000)直接溢出成无穷大,所有概率都变成NaN。
解决办法是数值稳定化:从输入中减去最大值,再求Softmax。这个技巧在数学上完全等价,因为:
def softmax(a): c = np.max(a) exp_a = np.exp(a - c) sum_exp_a = np.sum(exp_a) return exp_a / sum_exp_a为什么要减去最大值?因为a - c的最大值变成了0,np.exp(0)等于1,其他项都在0到1之间,求和不会溢出。我最初看到这个操作时觉得“多此一举”,直到自己跑模型输出NaN才发现,数值稳定性不是可选项,是必须项。
4.5 两层神经网络的完整算术流程
鱼书第四章实现两层神经网络时,前向传播的算术过程大概是这样的:
- 输入
X形状为(batch_size, 784),第一层权重W1形状为(784, hidden_size),np.dot(X, W1)得到(batch_size, hidden_size) - 加上偏置
b1形状为(hidden_size,),广播后得到加权和 - 通过激活函数(比如ReLU或Sigmoid),得到第一层输出
- 第一层输出和权重
W2做矩阵乘法,加上偏置b2,得到最终输出
我在实现这个流程时,最深的体会是:每个中间变量的形状变化都应该心里有数。为此我写代码时会专门加一行注释:
# X: (100, 784) -> W1@X.T -> (100, hidden) -> ReLU -> (100, hidden) -> W2@ -> (100, 10)这样即使睡一觉起来再看代码,逻辑也是通的。
5. 环境搭建与常见算术计算报错排查
5.1 Python和NumPy环境的快速搭建
如果你还在为“python怎么装”“numpy库装不上”纠结,不用急,我实测下来最省心的路径是:
- 官网下载Python,我建议3.8以上的稳定版本,鱼书代码在Python 3.8到3.12上跑都没问题
- 安装完打开命令行,输入
python --version确认安装成功 - 用pip安装NumPy:
pip install numpy - 装完之后,输入
python -c "import numpy as np; print(np.__version__)"测试
如果pip下载慢,除了上面提到的清华镜像源,还可以用-i参数指定阿里云镜像:
pip install numpy -i https://mirrors.aliyun.com/pypi/simple/Python环境配置的坑往往不在安装本身,而在命令行里调用了多个版本的Python。比如机器上装了Anaconda,又装了官方Python,输入python时可能调用了不同路径下的解释器。这时候pip和python可能不是同一个环境。我的经验是:在命令行里用where python(Windows)或which python(macOS/Linux)先确认自己到底在用什么,免得装了库却import不到。
5.2 最常见的算术计算报错与排查思路
我整理了几个我踩过的坑,做成一份速查表,你遇到了直接对照:
| 报错或异常 | 原因 | 排查方式 |
|---|---|---|
operands could not be broadcast together | 两个数组形状不匹配 | 打印双方的.shape,逐维对比广播规则 |
shape mismatch in numpy.dot | 矩阵内积维度不对应 | 确认前者列数等于后者行数 |
| 结果为NaN | 数值溢出或log(0) | 检查是否有极大值,加np.clip或delta保护 |
| 结果为全部0或全部1 | 使用//代替了/ | 检查除法运算符 |
| 数组元素全是整数 | dtype是int,运算结果被截断 | 用astype(np.float64)或初始化时传浮点数 |
| 训练损失不下降 | 学习率过大或梯度方向错误 | 打印梯度范数,调小学习率 |
广播报错是最常出现的。我的排查步骤基本固定:先打印两个数组的shape,再回忆一下广播三条规则,最后尝试用reshape或np.expand_dims调整形状。有一次我死活看不出来问题在哪,后来把两个数组的具体值也打出来了,才发现原来前一个数组的形状不是我以为的样子——它在某个环节被sum给降维了。所以,打印中间变量,永远是排查算术问题最可靠的手段。
5.3 浮点数精度与比较的隐藏陷阱
算术计算还有一个容易忽略的点:浮点数不能直接用==比较。比如你算了一个损失值,想判断它是否等于0,用if loss == 0很可能永远为False,因为浮点数在运算过程中积累了微小误差。正确做法是用一个很小的阈值判断:
if abs(loss) < 1e-7: print("损失已接近0")这不是鱼书特有的问题,而是所有数值计算共通的注意事项。我在处理梯度时也遇到过类似的情况:理论上某个梯度应该是0,打印出来却是-2.4e-17。这种细微的“脏数据”在后续累积计算中可能放大成明显的偏差。遇到这种情况,我的处理方式是适当做数值裁剪或四舍五入,比如np.round(grad, 6),让数据干净一点。
5.4 性能优化:避免Python原生循环
鱼书的代码样本不算大,但如果把训练数据扩到几十万条,Python原生循环会成为性能瓶颈。NumPy的向量化运算之所以快,是因为它在底层用C语言连续内存操作,避免了Python解释器逐元素调用的开销。我实测过:对100万个元素逐项求平方,用Python循环耗时大约0.3秒,用NumPy向量化运算耗时0.005秒,差距接近60倍。
因此,凡是能用数组运算解决的问题,就不要写循环。典型的例子是批量数据的损失计算:
loss = -np.sum(t * np.log(y + 1e-7)) / batch_size这行代码一次搞定所有样本的损失,等价于几十行循环代码。鱼书之所以值得精读,很大程度也是因为它不断在教你“怎么用算术思维而不是循环思维”去解决问题。
6. 实操心得与延伸思考
6.1 把算术计算当作调试工具
很多人把算术计算看成“跑数据前必须写的模板代码”,但我后来发现,算术计算本身就是调试工具。比如你在反向传播实现完后,不确定对不对,最直接的办法就是实现数值梯度,把数值梯度和解析梯度对比:
np.max(np.abs(grad_numerical - grad_analytic))如果最大误差在1e-5量级,基本说明反传写对了。这种验证方式本质上就是把“算术计算”变成判断标准。鱼书里虽然没有大篇幅讲这个技巧,但顺着书里的代码往下推,自然就会用到。
我在自己项目里,凡是实现新网络结构,都会把这个梯度检查作为默认步骤。它虽然会拖慢一点点速度,但能在第一时间把隐藏在矩阵运算里的错误揪出来。不少初学者在反向传播出错后,白白训练了好几天模型,损失曲线时而下降时而飙升,就是没用这个检查手段。
6.2 数学公式和代码之间的翻译能力
读鱼书最核心的能力,是把数学公式翻译成Python算术表达式。这种翻译能力是练出来的。我的练习方法很简单:每看到一个公式,先用纸笔把公式里每一步涉及的运算拆开,再对照书里代码找对应的NumPy函数。
比如Softmax公式:
y_k = exp(a_k) / sum_j exp(a_j)翻译成Python,就是先算所有exp,再求和,最后逐元素除法。翻译的关键在于识别公式里的下标和求和范围对应代码里哪个轴,以及除法对应的是逐元素运算还是矩阵运算。翻译得多了,你会发现神经网络里90%的数学公式,最后的落点都是加减乘除、指数、对数、求和、均值这几板斧。
6.3 我在实践中踩过的最值得分享的一个坑
最后分享一个我印象最深的问题:实现三层网络时,中间层激活函数用Sigmoid,结果训练到第二个epoch损失就变成了NaN。我一开始以为是学习率太大,调到0.001还是NaN,又怀疑是数据没归一化,折腾了一个多小时。后来打印各层的输出,发现隐藏层的输出在某个样本上达到了-1000,再送进Sigmoid,np.exp直接溢出。
这个问题的根源,是我初始化权重时用了标准差为1的正态分布随机数,导致深层信号在层层累积下爆炸。后来我把权重初始化改为标准差为1 / sqrt(n)的Xavier初始化,损失曲线才稳定下降。这个经历让我意识到,算术计算不只是“会算”,还得在数值稳定性和初始化策略上有意识。
鱼书后续的章节也讲了ReLU比Sigmoid更不易饱和、权重初始化的意义,但如果你在最开始的环境搭建和算术计算阶段就被NaN卡住,后面的内容很难推进。所以,我强烈建议你按这篇文章的顺序,先把Python原生运算、NumPy数组运算、广播机制、Softmax稳定化这些点一个个实测过关,再往后读。每一步代码都亲自跑一遍,把中间变量的shape和值都打印出来,你才能真正吃透这本“鱼书”。