05 GPU 与 TPU
05 GPU 与 TPU
Section titled “05 GPU 与 TPU”scalling 的关键在于如何合理的利用资源,而了解 GPU 在设计推理等架构上也非常重要。
GPU 硬件
Section titled “GPU 硬件”CPU 是通用架构,有分支单元等,是严格的串行设备。GPU 是并行的设备,有非常多的小计算单元。

GPU 的核心由 Streaming Multiprocessor 构成,每一个都是独立的计算单元,可以做单独的运算工作。非常多的 SM 就可以执行并行的运算系统。

GPU 的存储也分成很多层级,就如同 CPU 的寄存器等,平时常说的 GPU 显存大小一般都是 Global Memory,在此之上还有 L1, L2, Shared Memory 设备,他们的访问速度比通用内存要快得多。Global Memory 一般都是外部的显存芯片,而缓存等都是在 GPU 核心内部的。
因此优化运行效率的最主要方法,就优化分组 Block 级别的内存,让他尽可能减少使用低速存储。
GPU 编程
Section titled “GPU 编程”GPU 内部执行内部存在三种主要的术语:Thread, Block, Warp
- Thread :Thread 遵循的规则是 SMIT 规则,也就是所有 Thread 必须执行的是同一个命令(但是他们的输入可以不同)
- Block:Block 是一组 Thread。它的特点是他们可以访问同个 Shared Memory
- Warp:属于调度的概念,指的是一组同时运行线程的数量总是一个 Warp 为单位,例如 32 Threads 一个 Warp。
TPU 是 GPU 的另一种发展路线,但是 TPU 是更加针对于机器学习的一种计算单元。
TPU 的的内存,计算单元等架构和 GPU 非常像,区别在于他们的大小,控制器等。TPU 的计算单元大小一般来说比 GPU 要少,但是他的矩阵运算单元会更大。因此 GPU 会有更高的灵活性,而 TPU 只适合去做大矩阵的运算。
运算速度优化
Section titled “运算速度优化”
该图展示了运算优化的方向。在斜坡区域,计算的瓶颈在于内存传输等。随着坐标的右移,单位数据上的运算量会逐步上升,因而计算单元的利用率也会越来越高,最后计算单元饱和,达到水平区域,也就是我们需要优化的最终目标区域。
CPU 非常擅长分支运算,但是 GPU 很不擅长。
GPU 会同时计算一个 Warp 的 Thread,如果出现了分支,就会让一部分计算空闲着(因为一组 Thread 同一时间只能执行一个命令,不属于这个分支的 Thread 只能等着)
这也是为什么 ReLU 的实验中不是使用 If 语句来实现,而是通过掩码等执行确定化的运算。
NV 花了很大的经历在这方面。降低运算精度,减少需要搬运的数据,从而降低内存的带宽瓶颈。也就是一个 FLPO 计算的时候,更高精度的数据需要搬运更多的数据。
低精度的实现在 GPU 实现中实现相对复杂。同样模型在设计上也要考虑很多,例如在指数运算上对精度就很高,如何确定不同参数的精度也是一个难点。
在 FP16 级别,大家几乎都在使用 BF16,在 FP8 领域出现了分歧(有的指数多,有的尾数多)
在 BlackWell 架构中出现了 MXFP8 架构,在同一组数据中可以使用不同的缩放数,可以应对同一个矩阵中激活值跨度非常大的问题。

当然这样做肯定会带来弊端。例如 MXFP8 中转置矩阵时候会很复杂,在实际的硬件中,GPU 会自动保存一个转置的副本。
在模型中使用量化可以提高计算方法,但同样也要考虑好那些层可以进行量化。
Operator fusion,也就是尽可能将多个运算一次性做好,而不是分多次完成来付出需要搬运数据的代价。
一般来说这种简单的基础的算子融合,编译器会帮你自动优化好,让一次访问 Global Memory 尽可能完成更多的运算。
例如有 3 个连续的 sigmoid,为了实现反向传播,需要保存每一层的激活值。
为了优化这个存储,可以直接把这个计算的结果丢弃,反过来重新计算这个数据。(回顾一下,我们很多情况下瓶颈不在运算而在于存储通信等)在计算量不大的情况下这是一个很好的 trade-off
Memory soalescing,在 DRAM 硬件结构设计中,通过 burst mode 可以允许你一次性访问一大块连续的数据(例如 128 字节)

假如说一个 Warp 中 4 个 Thread 可以同时访问一块数据,通过合并访问可以一次性为 4 个线程准备好数据。体现在实际应用中就如上图,如果需要每个 Thread 访问矩阵中的不同的行,如果他在内存中是连续的,那么可以一次访问直接读取整个行。
这个操作对于大规模矩阵的读取比较有效,这也是为什么需要考虑矩阵在内存中的排列方式。
Tiling,尽可能地让内存同时被处理。


参考上图的矩阵运算。鉴于矩阵本身就有分块的性质(可以把一个大的乘法拆分成多个块矩阵的运算),另外两个 的矩阵每一个元素都需要被访问 n 次。可以分块的矩阵从 Global Memory 读到 SM 中,然后在 SM 反复读写,在所有的计算完成后再一次写回 GM。
同样这也会带来一些问题,如果矩阵大小是 tiling size 整数倍那很好,如果不是可能就会导致 block 的空间浪费。同时结合合并访问,如整数次合并访问不能正好读完一个 tiling size,也会造成浪费。在浪费严重的情况下会导致性能倒退。
这些整数倍的巧合也是导致 GPU 的计算能力等并不是纯线性增长的,一个基础的规律就是参数最好设置为 2 的幂(最好是 32 的倍数)

现在我们可以去尝试阅读这张图,并可以理解一下其中为什么会有突然掉下去的性能。这是因为不同的点,分别有可能被 2, 8, 16, 32 … 整除,有的甚至是奇数。在到达 16 以后几乎不太会有更多的性能损失,已经基本满足 burst mode。
CS336 对一些细致的小点做了更详细的解释,基本都是一些小的数值累加导致运算所需的次数等增加,此处不再详细解释。
Flash Attention
Section titled “Flash Attention”Flash Attention 是一种节省内存的高效运算方式。

FlashAttention 中将 KQV分块, 执行了 Tiling 运算,来加速。难点在于 Softmax 的运算是全局性的,如何能分块执行?具体的实现就是 Online Softmax,这是一种随着数据实时更新的 Softmax 计算方法。同时他也做了重计算的操作等。