L0.3 数学与机器学习基础
三维坐标
layer: L0(基础)|level: Engineer|pillar: 训推框架本文不是一篇数学复习课。目标是用 AI Infra 工程师的视角重读线代/概率与 ML 基础——搞清楚「为什么线性代数决定了底层算子是 GEMM 还是 reduction」「损失/梯度/反向传播在硬件上对应哪些计算」「一次训练 step 的 forward→loss→backward→optimizer.step 如何映射到显存占用」,从而为后续 L3 的 ZeRO / 显存优化埋下认知伏笔。
学习目标
- 前置知识:读过 L0.1/L0.2(知道「三大支柱」与三堵物理墙);会写基础 Python,记得「矩阵乘」「求导」这两个词大概是什么即可。无需线代/概率高分。
- 学完产出:① 能把深度学习坍缩成 GEMM / reduction / element-wise 三类计算原语,并说清各自是 compute-bound 还是 memory-bound;② 会画一次训练 step 的 显存账本(参数 / 梯度 / 优化器状态 / 激活),并据此算出「一个 NB 模型要多少显存」;③ 理解为什么这张账本是 L3 ZeRO 切分的前置认知;④ 不靠任何框架,纯 NumPy 手写一个两层 MLP 跑通前向 + 反向,亲眼看 loss 收敛。
- 阅读姿势:盯住一句话——「框架只是把这些数学自动化了;你手推一遍,才真正读得懂 autograd 和 ZeRO 在干什么。」
背景与现状
很多人把「数学基础」当成算法岗的事,与 infra 无关。这是新人最大的误区。AI Infra 的所有硬核优化,本质都是在优化一组确定的数学运算在确定硬件上的执行效率——你不理解这组运算长什么样,就无从谈优化。
从 infra 视角看,整个深度学习可以坍缩成三类计算原语:
- 稠密矩阵乘(GEMM, General Matrix Multiply):来自线性代数的
Y = XW。Transformer 里的 QKV 投影、FFN、注意力打分全是 GEMM,占据训练/推理 70%–90% 的 FLOPS,是 Tensor Core 的主战场。 - 规约(Reduction):来自概率统计的
softmax、LayerNorm、sum/mean、All-Reduce梯度同步。规约是**带宽受限(memory-bound)**的,不像 GEMM 那样吃算力,却常常是性能瓶颈。 - 逐元素(Element-wise):激活函数
ReLU/GeLU、残差相加、dropout。几乎纯带宽受限,是 **kernel fusion(算子融合)**最爱合并的对象。
从产业演进看,对「数学 → 算子 → 硬件」这条链路的认知,经历了三个阶段:
- 2012–2017(AlexNet 后):框架(Caffe/TensorFlow)把数学包成黑盒,工程师只写
model.fit(),不关心底层算子。 - 2017–2020(Transformer 普及):模型变大后发现「同样的 FLOPS,GPU 利用率(MFU)只有 30%」,业界开始反向追问——是 GEMM 没打满 Tensor Core,还是 reduction/element-wise 把带宽吃光了?
- 2020 至今(大模型时代):算子级认知成为 infra 核心竞争力。FlashAttention 把注意力的 softmax(reduction)与 GEMM 融合、避免反复读写 HBM,单这一个 idea 就重塑了整个长序列训练——而它的根,就是「线代决定算子形态、概率决定规约模式」这条本文要讲透的链路。
业界信号:FlashAttention 论文核心贡献不是新数学,而是重排了已知数学运算的访存顺序;ZeRO 的核心也不是新算法,而是重新分配了训练 step 中各类张量的显存归属。这印证一个事实——大模型 infra 的护城河,建立在对「最基础的数学运算如何落到硬件」的深刻理解之上。
📅 2026 时效:FlashAttention 已演进到 FA3(2024 年中发布,截至 2026 年仍处 beta/RC、生产多数仍用 FA2)——它进一步榨干 Hopper(H100) 的异步特性(warp-specialization + TMA + ping-pong 调度)与 FP8 低精度,FP16 比 FA2 快 1.5–2×。但它的「内核思想」从未变:重排访存、减少 HBM 往返,依旧是本文这条「数学→算子→访存」链路的延续。
原理与架构
2.1 线性代数速通:为什么一切都是 GEMM
神经网络的一层,数学上就是一次仿射变换加非线性:
其中 x 是输入(形状 [B, d_in],B 为 batch),W 是权重([d_in, d_out]),xW 就是一次 GEMM。这条最简单的式子带出三个 infra 关键认知:
- batch 维 B 让矩阵-向量乘升级成矩阵-矩阵乘。单样本是 GEMV(向量乘,带宽受限、Tensor Core 吃不饱);攒一个 batch 就变 GEMM(算力受限、能打满 Tensor Core)。这就是为什么推理要做 batching——不是为了省事,是为了把硬件利用率从个位数拉到 50%+。
- 矩阵形状决定硬件效率