ARTICLE DETAIL

资讯详情

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

Triton、TPU 与 Pallas 实战指南:用 Python 编写 GPU Kernel,理解 TPU 硬件架构与加速器选型

Triton、TPU 与 Pallas 实战指南:用 Python 编写 GPU Kernel,理解 TPU 硬件架构与加速器选型 Triton、TPU 与 Pallas 实战指南用 Python 编写 GPU Kernel理解 TPU 硬件架构与加速器选型【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendiumCUDA C 功能强大却冗长而 Triton 让研究者能用纯 Python 编写高性能 GPU kernelTPU 则以完全不同的硬件哲学脉动阵列提供大规模训练的另一条路径。本文以本仓库《chapter 16 - SIMD and GPU programming》系列为上下文系统讲解 Triton 块级编程模型、以 Flash Attention 为案例的 kernel 融合与 online softmax、TPU 的 MXU/ICI/BF16 架构以及 JAX/Pallas 的低级 kernel 编写方式最终给出一张覆盖训练、推理、融合 kernel、跨平台部署等场景的选型决策表。读完你将能够独立编写、调优并验证自己的 Triton/Pallas kernel并基于负载特征在 GPU 与 TPU 之间做出合理选择。本文主体对应仓库中的 triton, TPUs and pallax注意文件名中的 pallax 即 JAX 的Pallaskernel 编写 API并深度融合了同章 GPU architecture and CUDA、why C and how ML frameworks work 以及 computer architecture 等文件的底层原理。跨平台 GPU 计算Vulkan/WebGPU见 vulkan compute and cross-platform GPU。一、为什么需要 Triton爬出 CUDA 的抽象阶梯前一章 GPU architecture and CUDA 教你用 CUDA C 直接控制线程你需要理解 warp 调度GPU 以 32 线程为一组同时执行同一条指令即 SIMT、共享内存 bank 冲突、寄存器压力、合并访问coalescing等大量底层细节才能写出快 kernel。例如 GPU architecture and CUDA 指出合并访问要求一个 warp 内连续的线程访问连续的内存地址否则一次 128 字节的事务只有极小比例被利用tiling 模式则是把数据块从慢速全局内存搬到快速共享内存再复用是所有高性能 GPU kernel 的核心技术。TritonOpenAI 开源是建立在 Python 之上的 GPU kernel 语言。它换了一种思考方式你不必再推理单个线程而是推理数据块block。线程映射、内存合并、共享内存管理与大量优化由 Triton 编译器自动完成。其核心价值在于CUDA C 要求开发者精通 warp 调度、共享内存 bank 冲突、寄存器压力、合并访问模式Triton 把上述绝大部分抽象掉让只懂 Python、不懂系统编程的 ML 研究者也能写出生产级 kernel性能上可达手写 CUDA 的 80%95%而开发成本只有其一小部分。这与 ML 框架的整体设计哲学一脉相承why C and how ML frameworks work 中明确指出PyTorch 2.0 的torch.compile在 GPU 后端正是用Triton编译算子并融合操作而 JAX 则通过 XLA 编译到 CUDA/PTX 或 TPU 指令。也就是说Triton 不只是独立工具它已经是 PyTorch 官方编译栈的内核引擎。二、你的第一个 Triton Kernel向量加法下面是从零编写的完整向量加法 kernel可直接在 Colab GPU 运行时执行import triton import triton.language as tl import torch triton.jit def add_kernel( x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr, # compile-time constant ): # Each program instance processes one block of BLOCK_SIZE elements pid tl.program_id(axis0) # which block am I? block_start pid * BLOCK_SIZE # Offsets for this block offsets block_start tl.arange(0, BLOCK_SIZE) # Mask to handle the case where n_elements is not a multiple of BLOCK_SIZE mask offsets n_elements # Load data (masked: out-of-bounds reads return 0) x tl.load(x_ptr offsets, maskmask) y tl.load(y_ptr offsets, maskmask) # Compute output x y # Store result tl.store(output_ptr offsets, output, maskmask) def add(x: torch.Tensor, y: torch.Tensor) - torch.Tensor: output torch.empty_like(x) n_elements output.numel() # Launch: one program per block grid lambda meta: (triton.cdiv(n_elements, meta[BLOCK_SIZE]),) add_kernelgrid return output # Usage x torch.randn(1000000, devicecuda) y torch.randn(1000000, devicecuda) z add(x, y)逐行拆解这段代码它与 CUDA 的关键差异一目了然维度CUDA C 的做法Triton 的做法线程管理显式声明threadIdx/blockIdx手动计算线程数无显式线程概念只关心**程序block**编号tl.program_id(axis0)数据视图每个线程处理一个标量tl.arange(0, BLOCK_SIZE)直接生成整个 block 的偏移向量其上所有运算隐式向量化边界处理需要手写if (idx n)的标量清理循环mask掩码完成边界判断越界读返回 0类似 AVX-512 的掩码寄存器见 x86 and AVX内存访问手动保证合并访问tl.load/tl.store自动处理合并编译方式nvcc离线编译triton.jit首次调用时编译为 PTXGPU 汇编并缓存之后直接复用启动时grid lambda meta: (triton.cdiv(n_elements, meta[BLOCK_SIZE]),)表示每 BLOCK_SIZE 个元素启动一个程序实例BLOCK_SIZE: tl.constexpr声明该参数为编译期常量编译器可以据此做循环展开、寄存器分配等静态优化这也是 Triton 性能的关键机制之一——meta 字典中的值会在编译期嵌入 kernel。三、Triton Softmax Kernel一次 Kernel 融合的教学样板Softmax 是绝佳的 Triton 教学案例因为它天然需要多次遍历数据求最大值 → 减 max 取 exp → 求和 → 归一化且各趟之间共享同一份数据非常适合把数据留在 SRAM共享内存中复用triton.jit def softmax_kernel( output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, BLOCK_SIZE: tl.constexpr, ): # Each program handles one row row_idx tl.program_id(0) row_start input_ptr row_idx * input_row_stride # Load the row col_offsets tl.arange(0, BLOCK_SIZE) mask col_offsets n_cols row tl.load(row_start col_offsets, maskmask, other-float(inf)) # Softmax: max for numerical stability, then exp, then normalise row_max tl.max(row, axis0) numerator tl.exp(row - row_max) denominator tl.sum(numerator, axis0) softmax_output numerator / denominator # Store result output_start output_ptr row_idx * output_row_stride tl.store(output_start col_offsets, softmax_output, maskmask)几个值得注意的细节other-float(inf)掩码加载时越界位置填充-inf这样tl.max求出的行最大值不受填充位置影响——这是数值技巧与掩码机制结合的精妙之处整行数据只被加载一次max/exp/sum/divide 全部在寄存器与 SRAM 中完成中间结果不落全局内存反观 PyTorch 的F.softmax(x, dim-1)会启动3 个独立 kernelmax、exp-and-sum、divide每个都从全局内存读一次、写一次。这正是kernel 融合kernel fusion的意义内存访问次数从 6 次3 读 3 写降为 2 次1 读 1 写。对于 memory-bound 算子bias add、ReLU、layer norm 等内存流量往往就是执行时间的瓶颈融合直接把这部分流量抹掉。这也是为什么 GPU architecture and CUDA 强调PyTorch 的torch.compile和 Triton 让融合自动或近乎零成本地发生。实践中自定义 Triton kernel 相比 PyTorch 内置算子常有2–4 倍的加速尤其在算子融合场景这正是其价值所在。四、Triton 自动调优让编译器替你搜索最优配置Triton 内置auto-tuning在一次启动时对多组配置逐一基准测试选择实测最快的组合triton.autotune( configs[ triton.Config({BLOCK_SIZE_M: 128, BLOCK_SIZE_N: 128, BLOCK_SIZE_K: 32}), triton.Config({BLOCK_SIZE_M: 64, BLOCK_SIZE_N: 256, BLOCK_SIZE_K: 32}), triton.Config({BLOCK_SIZE_M: 256, BLOCK_SIZE_N: 64, BLOCK_SIZE_K: 64}), ], key[M, N, K], # re-tune when these change ) triton.jit def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K, ...): ...机制与使用要点configs中的每个triton.Config是一组编译期常量典型如矩阵分块尺寸 BLOCK_SIZE_M / BLOCK_SIZE_N / BLOCK_SIZE_K还可能包含num_warps、num_stages等调度参数key[M, N, K]声明当这些参数变化时才需要重新调优——Triton 会把 (M, N, K) 作为缓存键相同形状的矩阵直接复用已选出的最优配置避免每次启动都重新 benchmark调优在真实硬件上进行因此会自动适配不同 GPU 架构、矩阵维度与内存布局。最优分块尺寸与 GPU 架构如 Ampere A100、Hopper H100、SRAM 容量、矩阵形状强相关手调极难一次命中auto-tuning 恰好替你完成这轮穷举。五、Triton vs CUDA选型对比与边界| | Triton | CUDA C | |--|--------|--------| | 语言 | Python | C/C | | 抽象层次 | Block 级 | Thread 级 | | 开发速度 | 快每个 kernel 约 10–50 行 | 慢100–500 行 | | 性能上限 | 手写 CUDA 的约 80–95% | 100%完全掌控硬件 | | 共享内存 | 自动管理 | 手动管理 | | 内存合并 | 自动 | 手动 | | Warp 级原语 | 有限 | 完整shuffle、vote 等 | | 硬件支持 | 仅 NVIDIAAMD 为实验性支持 | 仅 NVIDIA |优先使用 Triton融合算子、自定义注意力模式如各类线性注意力/稀疏注意力变体、激活函数、以及绝大多数 ML 研究场景的 kernel 需求——你追求的是快速迭代 足够好的性能。优先使用 CUDA C当需要压榨最后 5–20% 的性能、需要 warp 级原语如__shfl_sync做跨线程归约、面对复杂的数据相关并行或 Triton 无法表达你的模式时。注意两张表都说明 Triton 目前主要面向 NVIDIA 生态AMD 支持处于实验阶段。若你的目标平台是移动端、浏览器或其他厂商硬件请参考 vulkan compute and cross-platform GPU 的 Vulkan/WebGPU 方案。六、案例研究Flash Attention——近年最具影响力的自定义 KernelFlash AttentionDao et al., 2022是近年 ML 领域最有影响力的自定义 kernel它以 $O(n)$ 内存复杂度代替标准注意力的 $O(n^2)$直接催生了超长上下文模型。6.1 问题注意力矩阵根本装不进显存标准注意力计算 $\text{softmax}(QK^T / \sqrt{d}) \cdot V$其中 $QK^T$ 是 $n \times n$ 的中间矩阵。当序列长度 $n 128K$ 时该矩阵为 $128K \times 128K \times 4$ 字节 ≈64 GB远超任何单卡显存。经典实现必须把它分块写入 HBM高带宽全局内存再读回来做 softmax产生巨量内存往返。6.2 洞察永远不物化完整的 $n \times n$ 矩阵核心洞察是按 tile 分块计算注意力。一次只加载一块 $Q$ 和一块 $K$在 SRAM 内算部分注意力分数、累积再移到下一个块。完整的 $n \times n$ 矩阵从未被物化——任一时刻 SRAM 中只存在一个 tile。这正是 GPU architecture and CUDA 介绍的tiling 模式在注意力场景的极致应用。6.3 难点Online Softmax棘手之处在于 softmax 的数值稳定性要求整行的最大值——而分块计算时你并不知道后面块里会不会出现更大的值。Flash Attention 用online softmax技巧解决维护一个运行中的最大值每当发现新的更大值时用缩放因子 rescale 之前已经算好的结果。这样 softmax 可以一块一块地增量完成数学上完全等价。6.4 算法伪代码For each block of Q rows: For each block of K columns: 1. Load Q_block from HBM to SRAM 2. Load K_block from HBM to SRAM 3. Compute S_block Q_block K_block.T (in SRAM) 4. Update running max, rescale previous results 5. Compute exp(S_block - running_max) 6. Update running sum and output accumulator Load V_block and compute final output Write output block back to HBM6.5 为什么快把内层循环钉死在 SRAM 里关键性能来源是数据复用内层循环完全在 SRAM 中进行HBM 只在加载 Q/K/V 块和写回最终输出时才被访问。SRAM 的访问延迟比 HBM 快约100 倍数据复用因子与 SRAM 容量成正比——tile 越大、复用越多性能越好但受 SRAM 容量与 bank 冲突约束这正是 auto-tuning 中num_warps、分块尺寸要调的原因。Flash Attention 同时有 Triton 与 CUDA C 两个实现CUDA 版本效率更高约快 10%但 Triton 版本可读性、可修改性远胜这对研究新型注意力变体linear attention、sparse attention 等至关重要——你可以快速改一版然后跑 benchmark而不是花数周改 C。七、TPU 架构与 GPU 截然不同的硬件哲学TPUTensor Processing Unit是 Google 自研的 ML 专用加速器其设计哲学与 GPU 大相径庭。7.1 脉动阵列Systolic Array与 MXUTPU 的核心计算单元是Matrix Multiply UnitMXU一个 128×128 或 256×256 的脉动阵列数据从阵列边缘流入在乘累加MAC单元网格中逐级传播每个单元完成一次乘加后把结果传给下一个单元。与 GPU 调度成千上万个独立线程不同脉动阵列是单一的确定性数据流没有线程调度、没有 warp 发散、没有分支预测这种简洁性让 MXU 在矩阵乘法上拥有极高的能效比每瓦 FLOPS代价是灵活性差非矩阵运算必须绕道向量单元vector unit速度显著慢于 MXU。对照 GPU architecture and CUDAGPU 的 Tensor Core 本质上是嵌入在通用 SM 里的专用矩阵乘法单元一条指令完成 4×4 矩阵乘 D A×B C而 TPU 的 MXU 把这一思想推到了极致——整个芯片的算力几乎都围绕矩阵乘法组织。7.2 HBM 与 ICIHBMTPU 使用与 GPU 相同的高带宽内存HBM。例如 TPU v5e 每芯片 16 GB HBM2eTPU v5p 每芯片 95 GB HBM2e具体容量以 Google 官方规格为准。ICIInter-Chip InterconnectTPU pod 通过自研高速网络连接数百颗 TPU 芯片。数据并行与模型并行见 distributed deep learning在 JAX 中原生支持跨 pod 的大规模并行训练因此成为 TPU 的核心卖点。相比 GPU 侧 NVLink单节点 8 卡 InfiniBand跨节点的两级网络ICI 的集成度与带宽在超大集群训练中通常更具优势。7.3 BFloat16为 ML 而生的数值格式TPU 是最早使用bfloat16的硬件。computer architecture 给出了精确定义bfloat16 为 1 8 7 16 位指数位与 float32 完全相同8 位只是尾数缩减到 7 位。这意味着与 float32 相同的指数范围 → 训练中不会溢出梯度值范围很大更少的尾数精度 → 对梯度更新足够相比 float165 位指数避免了小数值下溢与大数值上溢的两难。这个全指数、半精度的权衡对 ML 几乎是理想配置如今 PyTorch 的 BF16 混合精度训练distributed deep learning 提到的大模型训练标配正是这条技术路线的延伸。八、编程 TPUJAX 与 Pallas8.1 JAX XLA 编译管线TPU 不直接暴露 CUDA编程路径是JAX XLA你写 Python/JAX 代码jax.jit把它编译成 XLA 的中间表示 HLOHigh Level OperationsXLA 再把 HLO 编译为 TPU 专用指令。why C and how ML frameworks work 详细说明了这条管线jax.jit会 trace 你的函数、构建计算图、融合算子、消除冗余计算并优化内存布局然后针对目标后端CPU/GPU/TPU分别编译。全程无 CUDA、无 Cimport jax import jax.numpy as jnp jax.jit def matmul(a, b): return jnp.dot(a, b) # This runs on CPU, GPU, or TPU depending on the device a jnp.ones((1024, 1024)) b jnp.ones((1024, 1024)) c matmul(a, b)这份代码是设备无关的同一份matmul函数在 GPU 机器上编译为 CUDA/PTX在 TPU 机器上编译为 TPU 指令。这也是 TPU 生态以 JAX 为中心的根源——JAX 的函数式编程模型天然适配 XLA 的编译管线。8.2 PallasJAX 的 Triton 等价物Pallas是 JAX 的 kernel 编写 API定位等同于 Triton 之于 PyTorch让你在 Python 中编写低级 kernel由 XLA 编译到 GPU 或 TPUfrom jax.experimental import pallas as pl import jax.numpy as jnp def add_kernel(x_ref, y_ref, o_ref): o_ref[...] x_ref[...] y_ref[...] def add_pallas(x, y): return pl.pallas_call( add_kernel, out_shapejax.ShapeDtypeStruct(x.shape, x.dtype), grid(x.shape[0] // 128,), in_specs[pl.BlockSpec((128,), lambda i: (i,)), pl.BlockSpec((128,), lambda i: (i,))], out_specspl.BlockSpec((128,), lambda i: (i,)), )(x, y)关键概念解读x_ref/y_ref/o_ref是块引用block referencekernel 内通过切片语法o_ref[...] ...读写其语义与 Triton 的tl.load/tl.store对应pl.BlockSpec((128,), lambda i: (i,))声明每个程序处理的块形状128 元素以及如何从大张量中切出该块由 lambda 根据程序编号i计算偏移——这是 Pallas 版本的tl.arange 指针偏移grid(x.shape[0] // 128,)指定程序网格与 Triton 的grid lambda meta: (...)对应pl.pallas_call是编译入口out_shape声明输出张量的形状与 dtype。注意Pallas 目前比 Triton 更新、成熟度更低但它是唯一能为 TPU 编写自定义 kernel 的途径TPU 不支持 CUDA。如果你的项目跑在 GPU 上Pallas 同样可用但社区生态与资料量远不及 Triton。九、GPU vs TPU一张表看清权衡| | GPUNVIDIA | TPUGoogle | |--|-------------|--------------| | 可用性 | 任意云厂商、自建机房 | 仅 Google Cloud | | 编程方式 | CUDA C、Triton、PyTorch | JAX/XLA、Pallas | | 灵活性 | 通用计算 | 面向矩阵密集的 ML 优化 | | 峰值 matmul FLOPS | 很高Tensor Core | 很高MXU | | 非 matmul 算子 | 良好 | 较慢绕道向量单元不经 MXU | | 多芯片扩展 | NVLink8 GPU 内、InfiniBand | ICI数千 TPU集成度更高 | | 成本效率 | 有竞争力 | 大规模训练常更便宜 | | 生态 | 最大PyTorch、TensorFlow、JAX | 以 JAX 为中心 |选 GPU绝大多数 ML 负载、基于 PyTorch 的研究、推理服务、以及含显著非 matmul 计算的工作负载如大比例的自定义激活、复杂数据流控制。选 TPUGoogle Cloud 上数千芯片规模的 JAX 大规模训练、对成本敏感且以矩阵乘为主的负载如标准 Transformer 的大规模预训练。十、端到端选型Choosing the Right Tool工作负载最佳工具原因ML 训练PyTorchNVIDIA GPU CUDA/Triton生态最大、工具链最完善ML 训练JAX大规模TPU 或 NVIDIA GPUGoogle 规模下 TPU 更省成本GPU 更灵活自定义融合 kernelTritonPython或 CUDA CTriton 开发快CUDA 性能封顶JAX 自定义 kernelPallasTPU 唯一选择GPU 也可用跨平台推理Vulkanfile 07或 ONNX Runtime可在任意 GPU 厂商硬件上运行移动/边缘推理MetalApple、VulkanAndroid、NNAPI平台专用加速器浏览器推理WebGPUfile 07浏览器内唯一方案仅 CPU 推理ONNX Runtime AVX/NEON无需 GPU使用 SIMDfile 02、file 03新型硬件厂商专用 SDK每个加速器都有自己的工具链选型逻辑可以归纳为一句话先看框架PyTorch 生态优先 GPUJAX 大规模优先 TPU再看性能缺口最后 5–20% 性能才值得上 CUDA C最后看目标平台浏览器、移动端、CPU 各有专用方案。十一、动手练习建议使用 Colab GPU 运行时任务 1Triton 向量加法 vs PyTorch 内置加法import triton import triton.language as tl import torch import time triton.jit def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr): pid tl.program_id(0) offs pid * BLOCK tl.arange(0, BLOCK) mask offs n x tl.load(x_ptr offs, maskmask) y tl.load(y_ptr offs, maskmask) tl.store(out_ptr offs, x y, maskmask) n 10_000_000 x torch.randn(n, devicecuda) y torch.randn(n, devicecuda) # Triton out_triton torch.empty_like(x) grid lambda meta: (triton.cdiv(n, meta[BLOCK]),) add_kernelgrid # PyTorch out_torch x y # Verify correctness assert torch.allclose(out_triton, out_torch, atol1e-5) # Benchmark torch.cuda.synchronize() start time.time() for _ in range(1000): add_kernelgrid torch.cuda.synchronize() triton_time (time.time() - start) / 1000 start time.time() for _ in range(1000): out_torch x y torch.cuda.synchronize() torch_time (time.time() - start) / 1000 print(fTriton: {triton_time*1000:.3f} ms) print(fPyTorch: {torch_time*1000:.3f} ms) print(fRatio: {torch_time/triton_time:.2f}x)练习要点向量加法是**内存受限memory-bound**算子两者的差距主要来自 kernel 启动开销而非计算本身——你会观察到 Triton 与 PyTorch 非常接近这正是内存带宽瓶颈的体现。可以尝试把 BLOCK 从 256 调到 4096观察不同块大小对性能的影响理解为什么 auto-tuning 有价值。任务 2Triton 融合 kernel乘 加 ReLU 单趟完成import triton import triton.language as tl import torch import time triton.jit def fused_mul_add_relu_kernel(x_ptr, w_ptr, b_ptr, out_ptr, n, BLOCK: tl.constexpr): pid tl.program_id(0) offs pid * BLOCK tl.arange(0, BLOCK) mask offs n x tl.load(x_ptr offs, maskmask) w tl.load(w_ptr offs, maskmask) b tl.load(b_ptr offs, maskmask) result tl.maximum(x * w b, 0.0) # fused: mul add relu tl.store(out_ptr offs, result, maskmask) n 10_000_000 x torch.randn(n, devicecuda) w torch.randn(n, devicecuda) b torch.randn(n, devicecuda) # Fused (Triton) out_fused torch.empty_like(x) grid lambda meta: (triton.cdiv(n, meta[BLOCK]),) fused_mul_add_relu_kernelgrid # Unfused (PyTorch) out_unfused torch.relu(x * w b) assert torch.allclose(out_fused, out_unfused, atol1e-5) # Benchmark torch.cuda.synchronize() start time.time() for _ in range(1000): fused_mul_add_relu_kernelgrid torch.cuda.synchronize() fused_time (time.time() - start) / 1000 start time.time() for _ in range(1000): out_unfused torch.relu(x * w b) torch.cuda.synchronize() unfused_time (time.time() - start) / 1000 print(fFused (Triton): {fused_time*1000:.3f} ms) print(fUnfused (PyTorch): {unfused_time*1000:.3f} ms) print(fSpeedup: {unfused_time/fused_time:.2f}x)练习要点这是任务 1 的反面——PyTorch 的torch.relu(x * w b)会启动 3 个独立 kernel乘、加、ReLU每次都在 HBM 上往返读写Triton 版本一趟完成。这个任务你会看到明显的加速比它直观演示了第三节讨论的 kernel 融合原理。真实世界中的 LayerNorm、RMSNorm、注意力等自定义融合 kernel 正是同样的思路。任务 3验证 JAX/XLA 的自动算子融合import jax import jax.numpy as jnp import time def chain_ops(x): x x * 2.0 x x 1.0 x jnp.maximum(x, 0.0) # ReLU x x / jnp.sum(x) return x chain_jit jax.jit(chain_ops) x jax.random.normal(jax.random.PRNGKey(0), (10000, 1000)) # Warm up _ chain_jit(x) jax.block_until_ready(_) # Eager (each op is a separate kernel launch) start time.time() for _ in range(100): y chain_ops(x) jax.block_until_ready(y) eager_time (time.time() - start) / 100 # JIT (XLA fuses operations) start time.time() for _ in range(100): y chain_jit(x) jax.block_until_ready(y) jit_time (time.time() - start) / 100 print(fEager: {eager_time*1000:.2f} ms) print(fJIT: {jit_time*1000:.2f} ms) print(fSpeedup: {eager_time/jit_time:.1f}x (XLA fuses the 4 operations into 1 kernel))练习要点eager 模式下每个算子都是一次独立的 Python→C→kernel 往返4 个算子 4 次启动 4 次 HBM 往返jax.jit后 XLA 把 4 个操作融合为 1 个编译好的 kernelPython 完全脱离热循环。这正印证了 why C and how ML frameworks work 所描述的框架架构——这也是Python 是方向盘C/XLA/PTX 是引擎的最直观证据。十二、延伸阅读本文是《chapter 16 - SIMD and GPU Programming》系列的一部分建议按序阅读以获得完整脉络why C and how ML frameworks workPyTorch/JAX 的 C 后端架构理解torch.compile如何用 Triton 编译算子x86 and AVXAVX-512 掩码寄存器与 Triton mask 的对应关系GPU architecture and CUDAwarp/SIMT、合并访问、共享内存 tiling、Tensor Core 与混合精度——理解 Triton 自动化的底层对象vulkan compute and cross-platform GPU当目标平台不是 NVIDIA 时的跨平台 GPU 方案RISC-V and embedded systems嵌入式与开放指令集视角的硬件选择。对应数学与系统背景可参考 computer architecturebfloat16 定义与 distributed deep learning数据/模型并行与大规模训练系统。本仓库的 MCP Server见 README 与 mcp/src/index.ts可将整本手册作为知识库供 AI 助手检索方便在编写 kernel 时随时查阅这些章节。【免费下载链接】maths-cs-ai-compendiumBecome a cracked AI/ML researcher/engineer with this unconventional textbook covering maths, computing, and ML with intuition.项目地址: https://gitcode.com/GitHub_Trending/mat/maths-cs-ai-compendium创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表