Back to blog

Fine-Tuning and Serving Gemma 4 31B on Google Cloud TPU: A Technical Comparison with GPU Baselines

End-to-end TPU migration for Gemma 4 31B: 1.61× faster training at 2.12× lower cost, with 66% higher long-context inference throughput

Fine-Tuning and Serving Gemma 4 31B on Google Cloud TPU: A Technical Comparison with GPU Baselines

一、论文概述

项目内容
标题Fine-Tuning and Serving Gemma 4 31B on Google Cloud TPU: A Technical Comparison with GPU Baselines
作者Jatin Kishnani, Mayank Goel, Amit Singh, Pulkit Agrawal, Sairanjan Mishra
论文https://arxiv.org/abs/2605.25645
代码https://github.com/h2loop/gemma-tpu
模型https://huggingface.co/google/gemma-4-31B
发布2026-05-25 (v1→v3)
许可CC BY 4.0

二、核心思想

问题定义

LLM 的微调和推理主要在 GPU 平台上进行(PyTorch + HuggingFace TRL + FSDP),但 TPU 在大规模并行计算和互连带宽方面具有独特优势。将 GPU 原生的训练/推理管线迁移到 TPU 面临一系列工程挑战:

  • JAX/Tunix 与 PyTorch/TRL 的框架差异
  • Mesh 配置和 sharding annotation 修正
  • LoRA 模块命名约定变化(特别是 kv_einsum 融合)
  • Orbax checkpoint 到 safetensors 的自定义合并流程

解决方案概述

本文是第一篇端到端展示在 TPU 上微调和部署 Gemma 4 31B 的工作,系统记录了从 GPU 原生管线(PyTorch + HuggingFace TRL + FSDP)迁移到 JAX + Tunix/Qwix 栈所需的全部代码级适配,并提供了 TPU 与 GPU 平台的实证比较。

任务:Verilog 代码生成(CodeV-R1 数据集,10K 样本),在 NVlabs/verilog-eval benchmark(156 个 spec-to-RTL 问题)上评估。

三、技术架构

硬件配置对比

训练硬件:

规格TPU v5p-8GPU a3-highgpu-2g
加速器4× TPU v5p chips2× NVIDIA H100 80GB
每芯片 HBM102.8 GB80 GB
总 HBM411.2 GB160 GB
主机 RAM~355 GB~468 GB
互连ICINVLink
每小时费率$16.80/hr$22.12/hr

推理硬件:

规格TPU v6e-8 (Trillium)GPU a3-highgpu-2g
加速器8× TPU v6e chips2× NVIDIA H100 80GB
总 HBM250 GB160 GB
Tensor Parallelismtp=8tp=2
每小时费率$21.52/hr$22.12/hr

训练配置(两者相同)

参数值
优化器步数1,244 steps
Batch size8
序列长度4,096
数据集10K CodeV-R1 samples (1 epoch)
LoRArank=64, alpha=64
LR ScheduleAdamW cosine decay, peak 1e-4, 100 warmup

关键工程适配

1. Device Mesh 配置

JAX 需要显式 2D mesh。v5p-8 暴露 4 个 JAX devices(8 TensorCores / 2)。Mesh 为 (fsdp=1, tp=4)。约束条件是 num_global_key_value_heads=4——只有 tp=4 有效,不是 tp=8。

2. LoRA 模块命名映射

Tunix JAX 重命名了 HF PyTorch 模块,最关键的变化是 kv_einsum——将 K 和 V 投影融合为单个张量 (in, 2, n_kv_heads, head_dim)。

3. Sharding Annotation 修正

Qwix 的 LoRA 注入继承了原始权重的 PartitionSpec,但 LoRA 因子有不同的 rank。修复方案:将不匹配的 PartitionSpec 重置为完全复制,加上可整除性检查。

4. XLA Compiler Flags

xla_llvm_disable_expensive_passes 将编译时间从 ~8 分钟降到 ~3 分钟(牺牲 2-5% 运行时吞吐)。

5. Data Pipeline

将 HuggingFace DataLoader 替换为 Google 的 Grain 库,实现自定义 loss masking(仅 assistant tokens)和 reasoning strip(移除 xml... 块,仅保留 ```verilog 围栏)。

6. Gradient Checkpointing

Tunix 通过 nnx.remat 在 decoder 层粒度暴露。等价于 HF 的 activation_checkpointing=True。

Checkpoint 转换流程

Training Loss

自定义 orbax_to_peft.py 执行:

  1. 将 base model 加载到 JAX mesh
  2. 通过 Qwix 注入 LoRA + dummy inputs
  3. 恢复 Orbax checkpoint LoRA params
  4. 收集每模块的 lora_a/lora_b 对
  5. 应用 LoRA delta: W_merged = W_base + (alpha/r) * A*B^T
  6. 保存为单个 model.safetensors 文件

关键挑战:kv_einsum 映射到 TWO safetensors keys(独立的 K 和 V)。所有 deltas 转置以匹配 HuggingFace 约定(out, in)。

四、核心创新

创新点说明实验依据
完整 TPU 迁移记录文档化了 GPU→TPU 迁移的所有代码级适配6 个关键适配领域
Orbax→Safetensors 转换自定义 checkpoint 合并流程orbax_to_peft.py
TPU vs GPU 实证比较训练和推理的全面对比11 张表,6 张图
vLLM-TPU Docker 部署详述 TPU 上 vLLM 推理设置三大挑战:Docker-only、XLA 编译、HBM 分配

五、实验结果

训练性能对比

指标TPU v5p-8GPU 2×H100胜出
墙钟时间3.34 hr5.39 hrTPU (1.61× 更快)
吞吐量763 tok/s486 tok/sTPU (1.57× 更快)
每样本时间1.21 s1.94 sTPU (1.60× 更快)
最终 loss0.0722.258TPU
总成本$56.11$119.23TPU (2.12× 更便宜)
成本/百万 token$6.10$12.68TPU (2.08× 更便宜)

TPU 更快的原因

  1. 聚合 HBM 带宽:11.06 TB/s (TPU) vs 6.7 TB/s (GPU)
  2. ICI 互连:TPU 芯片间 900 GB/s 双向
  3. XLA fusion:前向+后向编译为单层融合 kernel
  4. GPU 无 torch.compile:因异构注意力层而禁用

评估结果(Verilog Code Generation)

Gradient Norm

指标TPU-trainedGPU-trained
pass@10.64100.6974
pass@50.79490.8141

GPU-trained 模型在评估质量上略优(归因于 loss 计算差异,而非硬件)。

Per-Problem Eval

推理性能对比

上下文长度指标TPU v6e-8GPU 2×H100胜出
短 (512/256)Peak throughput1,403 tok/s1,490 tok/sGPU (+6%)
短Median TTFT45 ms51 msTPU (1.1×)
中 (1024/512)Peak throughput1,404 tok/s1,387 tok/sTPU (+1%)
中TTFT @ QPS=449 ms58 msTPU (1.2×)
中QPS 饱和点~32~8TPU (4×)
长 (4096/512)Peak throughput1,206 tok/s728 tok/sTPU (+66%)
长 (4096/512)TTFT @ QPS=461 ms1,443 msTPU (23.6×)
长TPOT23.9 ms31.2 msTPU (1.3×)
超长 (8192/512)Peak throughput482 tok/s449 tok/sTPU (+7%)
超长TTFT @ QPS=41,013 ms7,202 msTPU (7.1×)
最大 (~16k/512)Peak throughput474 tok/s326 tok/sTPU (+45%)
最大TPOT13.5 ms19 msTPU (1.4×)

关键推理差异:

  • TPU 使用 fp8_e5m2 KV cache dtype(vs GPU bfloat16)→ 2× KV 容量
  • TPU 使用 tp=8(vs GPU tp=2)

推理成本对比

工作负载TPU $/1M tokGPU $/1M tok
短 (≤2048 tokens)$4.27$4.13
长 (4096 tokens)$4.95$8.44

总拥有成本 (TCO)

场景TPUGPU
训练 + 1hr 推理$77.63$141.35
训练 + 8hr 推理$228.27$296.19
训练 + 24hr 推理$572.59$650.11

推理成本对比图

Inference Cost

总结对比

Summary

六、TPU vs GPU 优劣势总结

TPU 优势

  • 训练速度 1.61× 更快
  • 训练成本 2.12× 更低
  • 长上下文推理吞吐 66% 更高
  • 长上下文 TTFT 23.6× 更快
  • 聚合 HBM 带宽更高(11.06 vs 6.7 TB/s)
  • ICI 互连带宽优势(900 GB/s)
  • fp8_e5m2 KV cache → 2× KV 容量

GPU 优势

  • 短上下文吞吐略高(+6%)
  • 生态成熟度(PyTorch/TRL/FSDP)
  • 评估质量略优(pass@1: 0.697 vs 0.641)
  • 更广泛的工具链支持

七、总结

核心贡献

  1. 首个 Gemma 4 31B TPU 端到端工作:从训练到推理的完整迁移记录
  2. GPU→TPU 代码级适配文档:mesh 配置、LoRA 命名、sharding 修正、gradient checkpoint、data pipeline、checkpoint 转换
  3. 全面的 TPU vs GPU 实证比较:训练 1.61× 更快/2.12× 更便宜,推理长上下文 66% 更高吞吐/23.6× 更快 TTFT
  4. vLLM-TPU 部署指南:Docker-only 部署、XLA 编译、HBM 分配三大挑战的解决方案
  5. 可复现代码:https://github.com/h2loop/gemma-tpu

技术影响

  • 为 GPU 原生 LLM 管线迁移到 TPU 提供了实用指南
  • 证明了 TPU 在长上下文推理中的显著优势(fp8 KV cache + tp=8)
  • 为云提供商的硬件选择提供了实证依据

八、参考资源

关键图片索引

图片说明文件名
Figure 1训练 loss 曲线training-loss.png
Figure 2梯度范数对比gradient-norm.png
Figure 3每问题 pass@1 对比eval-pass1.png
Figure 5推理成本对比inference-cost.png
Figure 6总结对比summary.png