纯 JAX 跑 Gemma 4:从 Turing 到 Ada,什么变了、什么没变
来源:dev.to — 2026-09-01
📋 概述
这是一份关于在跨代 NVIDIA GPU 上用纯 JAX 运行手写 Gemma 4 移植版的测量报告,聚焦"这就是 JAX"抽象在两处泄漏的地方——其中一处吃掉了 87% 的解码性能,而日志里没有任何红色报错。Turing(T4G)没有原生的 bf16 数据通路,XLA 会把它路由到 fp32,导致 54.1% 的解码时间耗在类型转换上,Tensor Core 全程零占用;而 Ada(L4)解决了存储与计算 dtype 匹配问题后吞吐提升 3.7 倍。作者强调 dtype 策略必须读取设备计算能力、Pallas 作为 API 可移植但作为内存模型不可移植、以及 KV 环缓存填充驱逐 bug 会返回 200 OK 和 "The The The The"。
🔑 核心要点
- 错误地选 compute dtype 不会报错,而是静默模拟:pre-Ampere 上 bf16 被路由到 fp32,解码消失
- dtype 策略不读配置文件,而是读取设备实时计算能力:COMPUTE_DTYPE = float16 if IS_PRE_AMPERE else bfloat16
- Turing 上 54.1% 解码耗在 dtype 转换、32.8% 耗在 fp32 gemvx,Tensor Core 0% 占用
- Pallas 作为 API 可移植,作为内存模型不可移植——两代卡上都跑不了快速路径
- KV 环缓存填充驱逐 bug 不崩溃、不产生 NaN,而是返回 HTTP 200 和 The The The The
- 修复存储与计算 dtype 匹配后吞吐提升约 3.7 倍,机器终于贴近带宽上限
💡 金句
最可怕的 bug 全都返回了成功。脚本里最吓人的 bug 都是静默返回 200 OK 的。
👍 0
👎 0
← 返回 dev.to 首页