ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

斯坦福大学 CS336 Lecture 05 GPU Principles and Distributed Training Fundamentals

斯坦福大学 CS336 Lecture 05 GPU Principles and Distributed Training Fundamentals 1. Goals①. Having a thorough understanding of the working principle and operation mechanism of GPU②. Knowing how to make fast algorithms like FlashAttention2. GPU Architecture2.1 Why ParallelismDennard Scaling随着晶体管变得越来越小它们的功率密度保持不变因此功率的使用与面积成比例Moores Law集成电路芯片上所集成的电路的数目每隔18个月就翻一番以蓝线所示尽管芯片的晶体管密度确实仍在持续增长单线程性能的增长曲线已明显趋于平缓。这意味着我们无法依赖于计算速度的提升为了获得更高的效率必须引入并行化。2.2 Difference between GPU and CPUCPU由于线程数量有限单线程运行速度极高。针对低延迟进行优化GPU由数量极大的计算单元构成注重并行运算。针对高数据吞吐进行优化。例如 T1 到 T4 任务序列即使 GPU 单个任务延迟可能更高但仍先完成序列。2.3 Layout of GPU基本架构就是 GPU - SM - SP/CUDA Core 以及现代 GPU 配置的 Tensor Core 以及 RT Core 处于 SM 中SMStreaming Multiprocessor基本运算单元。Triton 等工具的作用层级CUDA Core / SP Streaming Processor接受相同指令 将其应用于不同的数据Physical Layout MATTERS当数据吞吐量如此之大时内存与计算单元之间的物理距离开始产生重大影响。内存单元距离 SM 越近其访问速度就越快。Register、L1 Cache 以及 shared memory 就位于 SM 内部。而 L2 Cache 在 SM 外部仍在 GPU 芯片上。 VRAM 在芯片旁故数据读取速度逐渐降低。2.4 Execution Model of GPU为了实际编写针对 GPU 的高性能代码需要深入理解 GPU 的实际执行模型工作机制从线程块blocks、线程束warps以及线程threads三个层级进行深入考虑其粒度依次细化ThreadCUDA 中最小的执行单元每个 Thread 独享一组 Register。SM 里有一大块“寄存器文件Register File”里面包含成千上万个物理寄存器。Warp特定个 Thread 组成一个 WrapNVIDIA 为32AMD 为64。一个 Warp 里的所有 Thread 同时执行同一条指令但处理不同的数据SIMTSingle Instruction Multiple Threads。GPU硬件层面的调度单元。SM 里的调度器不是一个一个 Thread 去调度的而是以 Warp 为单位去调度。Block由程序员定义的一组 Thread数量可以自定。同一个 Block 里的所有 Thread被分配到同一个 SM上执行。一个 Block 上的所有 Thread 共享 shared memory。L1 Cache 是 SM 级别的资源它被该 SM 上同时执行的所有 Block可能不止一个共享。一个 Block 里的 Thread 数量会被硬件自动计算成 Warp 的整数倍。比如 33 个 Thread硬件会分配 2 个 Warp64 个 Thread第二个 Warp 里只有 1 个 Thread 在干活其余 31 个闲置。所以Block 里的 Thread 数量最好是 32 的倍数。2.5 About TPUsbased on Tensor Core2.6 Strengths of the GPU Model①. GPU 具有极好的延展性为了获得更高的速率不需要通过提高时钟效率只需增加更多 SM 单元即可。②. SIMT 架构适合处理矩阵运算等基础操作场景③. 线程都极其轻量化可以随时被暂停或重启蓝线GPU 与主机之间的数据传输带宽绿线全局内存的访问速度灰线计算速度当前性能瓶颈已转向内存带宽因为内存性能的提升速度远跟不上计算单元的演进这种发展趋势在未来仍将持续3. How to Make GPUs Go Faster横坐标是参与乘法运算的矩阵的维度大小纵轴则代表每秒执行的计算操作次数可以理解为硬件利用率。可以观察到随着矩阵维度变大硬件利用率不断增大。类似于渲染 Draw Call抵销任务调度等操作带来的额外开销但横纵坐标变化曲线表现出奇怪的波形需要探究其原因。上图的波形类似于屋顶线模型Roofline Model横轴为运算强度衡量的是“每读入一个字节的数据能进行多少次计算”。运算强度越高说明计算越密集。纵轴为计算性能。当我们考察性能时会发现存在两种典型状态第一种是内存瓶颈状态对应上图中左侧曲线部分大部分时间在等待数据从显存搬运到SM第二种是计算瓶颈状态对应图中右侧曲线部分计算非常密集已经把GPU的计算单元跑满了接近硬件的理论峰值算力。优化的最终目标处于右侧区域充分释放计算单元的效率3.1 Control Divergence (not a memory issue)由于 GPU 采用 SIMT 架构所以一个 warp 中的所有 thread 会执行同样的命令只是处理不同的数据。一个 warp 内部的条件语句会具有很强的破坏性会强制暂停所有未执行主控制流的线程破坏了 thread 的并行价值。所以在大规模并行计算单元中应该尽量避免使用条件判断语句3.2 Low Precision Computation对一个 n 维向量做一次逐元素的 ReLU : x max(0, x) 操作。若数据采用 fp32则每进行一次 flop 运算0 与 x 的比较需要进行读写操作各一次读入 x 输出结果 ReLU(x) 共两次 fp32为 8 bytes。计算强度为 8 bytes / FLOP如果采用 fp16则计算强度为 4 bytes / FLOP。在保持算法和硬件不变的情况下相当于免费获得了双倍的内存带宽。上图是在混合精度训练过程中哪些操作适合用什么精度数据。3.3 Operation Fusion在Lecture 04 4.3.2 DeepSeek v3 引入的 MLA 中优化 kv-cache 的过程中就提到了 Operation Fusion。3.4 Recomputation我去简单的思路极致的提升用富裕的计算资源换取紧缺的内存带宽其中3.5 Memory Coalescing and DRAM当读取某个内存地址时系统会一次性返回整块连续的内存数据而不只是单一数据这种机制称为突发传输模式burst mode连续的整块内存称为突发传输块burst section。其中原因是因为 DRAM 在读取数据时要经历一个昂贵的阶段称为行激活将电荷传递到放大器Sense Amplifier上将电信号转化成可读的 0/1 信号。之后的列读取和Burst Transfer阶段耗时都很短。类似于 Draw Call 和本章最开始的图所示矩阵越大硬件利用率越高是一样的。都是耗时很短的任务需要一个前置的代价很大的任务为了节省时间把多个小任务打包所以为了充分利用这种特性要精心设计内存访问模式。于是引出了内存合并Memory Coalescing当同一个 Warp 里的 32 个线程同时访问全局内存显存时如果它们访问的地址是连续的、对齐的硬件就会把这些零散的访问请求合并成一次或少数几次大的内存事务从而极大地提高内存带宽利用率。以上图所示若采用内存合并模式进行数据读取比对数据进行随机读取数据吞吐量为后者的四倍。以上图所示若要一系列 Threads 读取这个矩阵有两种方式一种是横向逐列读取一种是纵向逐行读取。永远记得 Threads 之间执行同样的命令横向读取时四个 Thread 在 T0 时刻会分别读取到M_0_0M_1_0M_2_0M_3_0无法实现 Memory Coalescing。而纵向可以实现极大提升读取效率。3.6 Tiling其实原理就是矩阵的分块运算一个最简单的 W 维方阵相乘一个 Thread 计算一个结果。如果不采用 Tiling计算整个方阵需要的总读取次数为 2 * W * W * W 次计算一个结果需要 2W共有 W^2 个数据换一个角度解释方阵 M 中的一个元素要与 方阵N 中的 W 个元素相乘方阵M、N共有 2 W^2 个元素。但如果采用 Tiling例如将一个 N 维方阵切成若干个 T 维方阵而划分后每次移动的最小单位从一个数字变成了一个 Tile。观察 M(0,0)它在 Tile_M(0,0) 中它需要与 W/T 个 Tile_N 进行计算故需要从 Global_Memory 中移到 shared_memory 中 W/T 次故两个方阵中所有元素一共需要 2 * W^2 * W / T 次耗时大的读取。而 Tile 内部的计算数据从 shared_memory 到计算单元读取效率高。所以只要 shared_memory 足够大能够容纳 tiling 后的矩阵块就可以将全局内存的数据传输总量降低至原来的 1/T。Complexity1.如何合理安排 Tile 大小a. 实现内存合并3.5b. 不能超过 shared_memory 能承载范围c. 尽量分配平均。一个 Tile 分配给一个 SM让 SM 间负载近似2. Memory Alignment : Interaction between tiling and burst section若 Tile 大小不是 burst section 大小的整数倍时就容易出现这种错位的情况。需要通过 padding 来调整 Tile大小。3.7 Summarize回到本章最开始的图重新分析一下。①. 总趋势类似于 Roofline前期未达到硬件的计算瓶颈后期已接近硬件的理论峰值算力②. 矩阵能否被 K 整除主要涉及 3.6 中所提到的 Tiling 和 Burst Section 是否对齐的问题。③. 断崖式下降以横坐标为 1792 到 1793 的骤降为例假设 Tile Size 为256 * 128则 Tile 数为 (1792 / 256) * (1792 / 128) 7 * 14 98。而当仅增长 1 到 1793 时Tile 数为 8 * 15 120。以 A100 为例它拥有 108 个 SM这会带来新的计算消耗。4. Flash Attention传统 Attention 的问题1.显存爆炸对于 S 和 P 中间结果需要存储在显存中当词序列较长时会出现 OOM 问题2.带宽瓶颈读取 QKV (3 * N * d) - 计算并写入 S (N * N) - 读取 S 并计算 Softmax (N * N) - 写入 P (N * N) - 读取 P 和 V 并计算 O ( V: N * d) - 写入 O (N * d)总访存量O(N² N×d)计算量 O(N²×d) 。当d较小时两者相近GPU算力被严重浪费4.1 Tiling for QKV Matrix4.2 Incremental Computation of Softmax稳定 Softmax 计算方法( 保证数据稳定不会溢出)传统 Softmax 无法在 Tile 中计算需要全部元素信息。数学原理很简单就是分块处理然后块与块之间比较做一次缩放就行。两者结合起来就是完整的 FlashAttention 前向传播的过程。
返回列表