1.4 运算与算子:排程单上的最小工单 本节摘要:算子(Operation)是计算图上最小的调度单元,一张排程单就是"某算子吃进若干张量、吐出若干张量"。本节梳理四大算子族的职责边界,讲清运算符重载、类型提升、归约运算的执行真相,并解释为什么"看起来一行代码"在图上可能登记了好几个节点。算子是排程的语言,说准这门语言才能读懂图的执行顺序。 本节能力清单 阅读完本节,你应当能够: 把一段张量代码手工翻译成"节点与边"的语言,数出图上有几个节点; 区分逐元素运算、矩阵运算、归约运算、形状运算四族算子的适用场景; 预判混合类型运算的类型提升结果,避免静默升精度带来的性能损失; 用归约类算子写出损失函数里最常见的求均值、求和、按轴最大值。 一行代码,几个节点 先做个小实验。
本节摘要:算子(Operation)是计算图上最小的调度单元,一张排程单就是"某算子吃进若干张量、吐出若干张量"。本节梳理四大算子族的职责边界,讲清运算符重载、类型提升、归约运算的执行真相,并解释为什么"看起来一行代码"在图上可能登记了好几个节点。算子是排程的语言,说准这门语言才能读懂图的执行顺序。
阅读完本节,你应当能够:
先做个小实验。下面一行看起来是一个运算:
import tensorflow as tf a = tf.constant([1.0, 2.0, 3.0]) b = tf.constant([4.0, 5.0, 6.0]) c = a + b * 2.0 # 一行 Python,两个工单 print(c) # 输出:tf.Tensor([ 9. 14. 19.], shape=(3,), dtype=float32)
图上实际登记了三个节点:常数 2.0 的源节点、一个 Mul(b 乘 2.0)、一个 Add(a 加上乘积)。+ 与 * 只是 Python 运算符重载,背后调用的分别是 tf.add 与 tf.multiply。数据流依赖决定了执行顺序:Mul 必须先于 Add 完成,因为后者的输入是前者的输出。读代码时在脑子里画依赖箭头,是排程室的基本功。
# 手工展开上面那行,节点就现形了 two = tf.constant(2.0) # 节点 1:常量源 mul = tf.multiply(b, two) # 节点 2:Mul,吃 b 与 two add = tf.add(a, mul) # 节点 3:Add,吃 a 与 mul tf.assert_all_finite(add) # 依赖检查通过,值已就绪 print(add) # 输出:tf.Tensor([ 9. 14. 19.], shape=(3,), dtype=float32)
排程单虽然成百上千张,按职责分就四族。逐元素运算对每个位置独立处理,是最容易并行的一族;矩阵运算是深度学习的主力,GPU 优化的重心;归约运算把多个元素压成一个或一串,是损失函数的常客;形状运算只重排不计算,成本近乎为零。
| 算子族 | 代表成员 | 典型用途 |
|---|---|---|
| 逐元素 | tf.add、tf.multiply、tf.nn.relu | 激活函数、逐位置修正 |
| 矩阵与线性代数 | tf.matmul、tf.transpose | 全连接层、注意力打分 |
| 归约 | tf.reduce_mean、tf.reduce_sum、tf.reduce_max | 损失聚合、统计特征 |
| 形状与拼接 | tf.reshape、tf.concat、tf.tile | 批次整理、多路特征合并 |
# 矩阵乘法:全连接层的数学本体 x = tf.constant([[1.0, 2.0]]) # (1, 2) 输入特征 w = tf.constant([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) # (2, 3) 权重 y = tf.matmul(x, w) # (1, 3) 输出 print(y) # 输出:tf.Tensor([[0.9 1.2 1.5]], shape=(1, 3), dtype=float32) # 手算第一列:1*0.1 + 2*0.4 = 0.9,对上了 # 逐元素与矩阵乘法的区别,形状约定完全不同 try: tf.multiply(x, w) except Exception as e: print(type(e).__name__) # multiply 要求形状可广播:(1,2) 对 (2,3) 不兼容 # 通常抛出 InvalidArgumentError
归约运算的关键参数是 axis,它指定"沿哪根轴压"。axis 的语义新手最容易绕晕,记住一句话:axis 是被消灭的那根轴。
m = tf.constant([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) # (2, 3) print(tf.reduce_sum(m)) # 不给 axis:全压成标量 # 输出:tf.Tensor(21.0, shape=(), dtype=float32) print(tf.reduce_sum(m, axis=0)) # 消灭轴 0(行方向压扁) # 输出:tf.Tensor([5. 7. 9.], shape=(3,), dtype=float32) print(tf.reduce_mean(m, axis=1)) # 消灭轴 1(列方向压扁) # 输出:tf.Tensor([2. 5.], shape=(2,), dtype=float32)
损失函数几乎都是"逐元素误差 + 归约"的组合,例如均方误差就是 tf.reduce_mean(tf.square(y_true - y_pred))——先逐元素相减平方,再沿全部元素求均值。
两个不同 dtype 的张量做运算,TensorFlow 会自动提升到兼容类型,而不是报错。方便,但有代价:
i = tf.constant([1, 2, 3], dtype=tf.int32) f = tf.constant([0.5, 0.5, 0.5], dtype=tf.float32) r = i + f # int32 静默提升为 float32 print(r.dtype) # 输出:<dtype: 'float32'> # 明确控制:该 cast 就 cast,别依赖提升 i2 = tf.cast(i, tf.float32) print((i2 * f).dtype) # 输出:<dtype: 'float32'>
隐患有两处:一是混合精度训练里,一个疏忽的整型运算可能把本该跑 float16 的部分拖回 float32;二是 Python 整数与张量运算时按"能装下就提升"的规则走,偶尔出现意料外的 float64。规范做法是数据入口统一 cast,后续运算不依赖隐式提升。
每个算子注册时都附带自己的求导规则,这是自动微分能工作的前提——反向传播不是魔法,而是每个工单都自带"怎么把上游梯度分回给下游"的说明书。用一个小例子验证:
x = tf.Variable(2.0) with tf.GradientTape() as tape: y = x * x # Mul 算子:说明书是 2x z = y * x # 再登记一次 Mul dz = tape.gradient(z, x) print(dz) # 输出:tf.Tensor(12.0, shape=(), dtype=float32) # z = x 的三次方,导数 3 倍 x 的平方,在 x=2 处即 12
乘方被拆成两次 Mul 登记,求导链也按两次 Mul 的说明书串联。1.5 节讲执行模式,第 5 章把这个求导机制彻底展开。
⚠️ 常见坑:
tf.matmul与tf.multiply只差两个字母,语义天差地别。全连接层写错成 multiply,轻则形状报错,重则广播成功后结果全错——审代码时对这两个名字保持一级警觉。
💡 关键直觉:报错信息里出现 "node" 或 "op" 字样时,说的就是图上的算子。把报错里的算子名回填到代码里定位节点,比逐行重读代码快。
下节看这些工单以什么节奏被执行——两种执行模式的调度差异。