h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
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
arXiv:2605.25645v1 [cs.DC] 25 May 2026
Sairanjan Mishra May 2026 Abstract We present the first end-to-end demonstration of fine-tuning and serving Google’s Gemma 4 31B model on TPU hardware, providing an empirical comparison of TPU and GPU platforms for large language model adaptation. Using LoRA on a Google TPU v5p-8 for training and TPU v6e-8 (Trillium) for inference, we document the full set of code-level adaptations required to port a GPU-native training recipe — built on PyTorch, HuggingFace TRL, and FSDP — to the JAX + Tunix/Qwix stack. These adaptations span mesh configuration, LoRA module naming conventions, sharding annotation corrections, gradient checkpointing, data pipeline restructuring, and a custom Orbax-to-safetensors checkpoint merging procedure. For inference, we detail the vLLM-TPU Docker setup necessary to serve Gemma 4 on v6e-8 and characterize the resulting latency and throughput profile. Compared with a 2×H100 GPU baseline under identical hyperparameters, TPU training completes 1.61× faster at 2.12× lower cost. Inference throughput is within 3% across platforms, while TPU achieves 2× lower time-to-first-token (235 ms vs. 475 ms). Together, the TPU configuration is 1.82× cheaper for a representative train-plus service workload. Our work removes a critical gap in the open tooling ecosystem and provides practitioners with a reproducible, production-ready recipe for Gemma 4 deployment on the TPU infrastructure.
Contents 1 Introduction
3
2 Hardware and Infrastructure 2.1 Training Hardware . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 2.2 Inference Hardware . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 2.3 TPU VM Environment . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
3 3 3 3
3 TPU Training Recipe 3.1 Framework Differences: PyTorch vs JAX . . . . . . . . . . . . . . . . . . . . . . . 3.2 Device Mesh Configuration . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 3.3 LoRA Module Naming: Tunix vs HuggingFace . . . . . . . . . . . . . . . . . . . 3.4 LoRA Sharding Annotation Fixes . . . . . . . . . . . . . . . . . . . . . . . . . . . 3.5 Gradient Checkpointing . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 3.6 XLA Compiler Flags . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 3.7 Optimizer: Optax AdamW . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 3.8 Data Pipeline: Grain vs HuggingFace DataLoader . . . . . . . . . . . . . . . . . 3.9 Checkpointing . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
4 4 4 4 5 5 5 6 6 6
1
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
4 Checkpoint Conversion: Orbax to Safetensors 4.1 Merge Process . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 4.2 Tensor Shape Mapping . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 4.3 Why Not Use PEFT Merge Directly? . . . . . . . . . . . . . . . . . . . . . . . . .
6 7 7 7
5 Training Results 5.1 Configuration . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 5.2 Loss Curves . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 5.3 Why is TPU Faster? . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 5.4 Training Cost . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
8 8 8 9 9
6 Evaluation Results 6.1 Benchmark . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 6.2 Per-Problem Analysis . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
10 10 10
7 TPU Inference: vLLM on v6e-8 7.1 Setup Challenges . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 7.1.1 Docker-Only Deployment . . . . . . . . . . . . . . . . . . . . . . . . . . . 7.1.2 XLA Compilation on First Startup . . . . . . . . . . . . . . . . . . . . . . 7.1.3 HBM Allocation . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 7.2 GPU vLLM Setup (for comparison) . . . . . . . . . . . . . . . . . . . . . . . . .
11 11 11 11 12 12
8 Inference Results and Analysis 8.1 Benchmark Configuration . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 8.2 Explaining the Short-Context GPU Advantage . . . . . . . . . . . . . . . . . . . 8.3 Explaining the Long-Context TPU Dominance . . . . . . . . . . . . . . . . . . . 8.4 Explaining the TTFT Advantage at Short Context . . . . . . . . . . . . . . . . . 8.5 Explaining Slightly Better GPU P99 Tail at Short Context . . . . . . . . . . . . 8.6 Inference Cost . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
12 12 13 13 13 14 14
9 Total Cost of Ownership
14
10 Summary 10.1 TPU Advantages . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 10.2 GPU Advantages . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 10.3 Conclusion . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . 10.4 Resources . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . . .
15 15 15 16 16
2
h2loop.ai — Technical Report
1
TPU vs GPU: Gemma 4 31B LoRA SFT
Introduction
Large language model fine-tuning is increasingly bottlenecked not by data availability but by compute cost and iteration speed. Google Cloud TPUs, while natively supported for training, require non-trivial porting effort compared to the mature GPU ecosystem (PyTorch, HuggingFace, FSDP). This report documents our end-to-end experience running Gemma 4 31B LoRA supervised fine-tuning on TPU v5p-8 and vLLM inference on TPU v6e-8 (Trillium), and compares both to an existing GPU baseline on 2×H100 80GB (GCP a3-highgpu-2g). The task is Verilog code generation: we fine-tune on the CodeV-R1 dataset [1] (10K samples of natural-language-to-Verilog pairs) and evaluate on NVlabs/verilog-eval [2] (156 spec-to-RTL problems, pass@1/pass@5 with iverilog simulation). The paper is structured as a walkthrough of the engineering changes required to make TPU work, followed by quantitative results and cost analysis.
2
Hardware and Infrastructure
2.1
Training Hardware Table 1: Training hardware configurations.
2.2
Spec
TPU v5p-8
GPU a3-highgpu-2g
Accelerator HBM per chip Total HBM Host RAM Interconnect GCP on-demand rate
4× TPU v5p chips 102.8 GB 411.2 GB ˜355 GB ICI (inter-chip) $16.80/hr
2× NVIDIA H100 80GB HBM3 80 GB 160 GB ˜468 GB NVLink $22.12/hr
Inference Hardware Table 2: Inference hardware configurations.
2.3
Spec
TPU v6e-8 (Trillium)
GPU a3-highgpu-2g
Accelerator HBM per chip Total HBM Tensor parallelism GCP on-demand rate
8× TPU v6e chips 31.25 GB 250 GB tp=8 $21.52/hr
2× NVIDIA H100 80GB HBM3 80 GB 160 GB tp=2 $22.12/hr
TPU VM Environment
Both TPU generations (v5p and v6e) ship with Python 3.10, but the Tunix framework requires Python 3.11+. The first setup step is therefore installing Python 3.11 via the deadsnakes PPA before anything else. The JAX TPU wheel is installed from a separate Google-hosted index: pip install " jax [ tpu ] " -f https :// storage . googleapis . com / jax - releases / libtpu_releases . html
A critical operational note: the TPU VM’s boot disk (97 GB) cannot hold the 62.5 GB Gemma 4 31B model download alongside the OS. Tunix stages the model in /tmp/models before loading 3
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
into HBM. On the v5p-8, the host has 355 GB of RAM-backed /dev/shm (tmpfs), which we redirect the staging to via TMPDIR=/dev/shm at launch time. Without this, the model download fills the boot disk and crashes mid-load. For the v6e-8 (Trillium) inference instance, we found that only the v2-alpha-tpuv6e runtime works correctly—the generic tpu-ubuntu2204-base runtime fails at TPU topology enumeration. Zone availability for v6e-8 was constrained; we used asia-northeast1-b as the only available zone at time of testing.
3
TPU Training Recipe
3.1
Framework Differences: PyTorch vs JAX
The GPU baseline uses PyTorch + HuggingFace TRL (SFTTrainer) with FSDP for model parallelism. The TPU recipe uses JAX + Tunix (PeftTrainer) with Qwix for LoRA injection. These are fundamentally different execution models: • GPU: Eager PyTorch with FSDP sharding. Each GPU sees its shard of the model; gradients are reduced across ranks by FSDP’s all-reduce. • TPU: XLA-compiled JAX with an explicit device mesh. The model is sharded via jax.sharding.PartitionSpec annotations; XLA generates the communication collectives automatically. The first training step incurs a 3–8 minute XLA compilation cost; subsequent steps execute the compiled graph. The key implication is that TPU training is less tolerant of dynamic shapes and control flow that exits traced functions. Every shape that varies across steps triggers a recompile.
3.2
Device Mesh Configuration
JAX requires an explicit 2D device mesh specifying how chips are arranged across FSDP and tensor-parallel (TP) dimensions. For the v5p-8 (4 chips): # v5p -8 has 4 chips ; tp =8 breaks on Gemma4 heads not divisible by 8 MESH_COUNTS = (1 , 4) # fsdp =1 , tp =4 mesh = jax . make_mesh ( MESH_COUNTS , ( " fsdp " , " tp " ) , axis_types =( jax . sharding . AxisType . Auto ,) * 2 , )
The tp = 8 configuration was attempted but failed for Gemma 4 31B because the model’s attention head count (32 heads, 16 KV heads) is not divisible by 8 for all projection shapes. tp=4 works correctly. In a hypothetical 16-chip configuration (v5p-16), the mesh would be (2, 8).
3.3
LoRA Module Naming: Tunix vs HuggingFace
The HuggingFace model has separate q proj, k proj, v proj, and o proj Linear layers. The vision-tower exclusion regex (^ (?!.*vision)) is required because the HF Gemma 4 checkpoint is multimodal and the vision tower contains similarly-named projections. Tunix’s JAX port of Gemma 4 uses different module names:
4
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
Table 3: LoRA target module name mapping between HF PyTorch and Tunix JAX. HF PyTorch (GPU)
Tunix JAX (TPU)
Notes
q proj k proj + v proj o proj gate proj up proj down proj
q einsum kv einsum attn vec einsum gate proj up proj down proj
Renamed; same weights Fused into a single tensor with dim-1 = {K,V} Renamed Same Same Same
The kv einsum fusion is the most significant difference: the Tunix model stores both K and V projections in a single weight tensor of shape (in, 2, n kv heads, head dim), while HuggingFace stores them as separate (n kv heads*head dim, in) matrices. This requires special handling during checkpoint merging (see Section 4). There is also no vision tower in the JAX port, so the exclusion regex is unnecessary: _LORA_MODULES = ( " .* q_einsum |.* kv_einsum |.* attn_vec_einsum " " |.* gate_proj |.* down_proj |.* up_proj " )
3.4
LoRA Sharding Annotation Fixes
Qwix injects LoRA parameters by tracing the model and inserting lora a/lora b tensors. However, it inherits the sharding annotations (PartitionSpec) from the original weight tensor— which has a different rank than the LoRA factors. For example, the kv einsum base weight has shape (in, 2, n kv heads, head dim) with a 4-element PartitionSpec, but the LoRA lora a has shape (in, rank) (rank 2). When the optimizer tries to shard optimizer states according to this mismatched spec, JAX raises an error. We fix this with a post-injection pass that resets any PartitionSpec whose rank doesn’t match the actual tensor rank to a fully-replicated spec. A second pass checks divisibility: if a tensor dimension is not divisible by the mesh size for the assigned axis, we replicate that axis to avoid uneven sharding. The same fix is applied to the optimizer state via a monkey-patched shard optimizer that filters the partition specs through the same divisibility check before applying jax.lax.with sharding constraint.
3.5
Gradient Checkpointing
On TPU, gradient checkpointing (called rematerialization in JAX) is applied at the decoder layer granularity. Tunix exposes this via nnx.remat. This patches the unbound method directly, so all decoder layer instances use the rematerialized forward pass. Without this, the 31B model’s activation memory during backward would exceed the v5p-8’s 411 GB HBM at the batch sizes used. The GPU equivalent is HuggingFace’s activation checkpointing=True in the FSDP config.
3.6
XLA Compiler Flags
The xla llvm disable expensive passes flag skips several LLVM optimization passes that are particularly slow for large Transformer graphs. This trades a small amount of runtime throughput (˜2–5%) for significantly faster XLA compilation (from ˜8 min to ˜3 min for the first step).
5
h2loop.ai — Technical Report
3.7
TPU vs GPU: Gemma 4 31B LoRA SFT
Optimizer: Optax AdamW
The GPU recipe uses PyTorch’s adamw torch optimizer via HuggingFace Trainer. On TPU, JAX has no built-in optimizer; instead we use Optax, Google’s composable gradient transformation library for JAX. This mirrors the GPU recipe exactly: linear warmup from 0 to peak LR over 100 steps. Followed by cosine decay to 0 over the remaining steps. The key differences from PyTorch AdamW: • Composability: Optax builds the optimizer as a chain of stateless gradient transformations (scale by adam, add decayed weights, scale by schedule). Each transformation has its own state pytree, which Orbax checkpoints independently. • No grad clipping by default: PyTorch’s HuggingFace Trainer applies max grad norm=1.0 automatically. On TPU we omit this (matching the measured gradient norms, which stay stable without clipping), though optax.clip by global norm(1.0) can be prepended to the chain if needed. • Schedule is a pure function: The LR at any step is computed by calling schedule(step) — useful for logging and the warmup-clamp we apply when running short subsample runs. • Optimizer state sharding: The Optax state (Adam moments) is sharded across the same device mesh as the model weights via shard optimizer, requiring the divisibility fix described in Section 3.4.
3.8
Data Pipeline: Grain vs HuggingFace DataLoader
The GPU recipe uses HuggingFace’s DataCollatorForCompletionOnlyLM with SFTTrainer. The TPU recipe uses Google’s grain library (a deterministic, shardable data loader). Key differences: 1. Loss masking: HuggingFace uses the {% generation %} chat template tag to identify assistant-turn tokens. Tunix’s tokenizer adapter doesn’t emit these markers, so we implement the mask by tokenizing user and model turns separately and computing offsets: 2. System prompt handling: The GPU recipe includes a system prompt in the HuggingFace chat template, but the Gemma template silently drops system messages—only user and assistant roles are rendered. The TPU recipe mirrors this exactly by emitting only user/model turns. 3. Reasoning strip: The CodeV-R1 dataset contains chain-of-thought reasoning in <think>...</think> blocks. We strip these and keep only the <answer> block (or the Verilog fence directly), reducing sequence lengths and focusing training on the final code output. 4. Drop vs truncate: Both recipes drop sequences longer than max seq len rather than truncating, to avoid training on incomplete Verilog modules. With max seq len=3072, 99.6% of the 10K-sample subset is retained (37 dropped as too long, 4 empty).
3.9
Checkpointing
The TPU recipe uses Orbax (orbax-checkpoint) with a GCS path for checkpoint storage. Checkpoints are saved every 100 optimizer steps, retaining the 3 most recent (matching the GPU recipe’s save steps=500, save total limit=3 with a finer cadence for TPU’s longer run). Checkpoints survive spot preemption because they are written directly to GCS rather than local disk.
4
Checkpoint Conversion: Orbax to Safetensors
After training, inference requires merging the LoRA adapters back into the base weights and saving a merged safetensors model. The GPU recipe uses HuggingFace PEFT’s built-
6
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
in merge and unload(). The TPU recipe requires a custom orbax to peft.py due to: 1. Orbax checkpoints store the full model tree including LoRA wrappers in a JAX-specific format (not safetensors). 2. The Tunix JAX model uses different weight names and tensor layouts than HuggingFace safetensors (which is what Tunix’s sampler ultimately loads from for inference). 3. The kv einsum fusion means one LoRA module maps to two safetensors keys.
4.1
Merge Process
The merging steps are: 1. Load the base model (from GCS safetensors) into JAX with the mesh. 2. Inject LoRA via Qwix (same structure as training) using dummy inputs. 3. Restore the Orbax checkpoint into the LoRA model (LoRA params only). 4. Iterate over all nnx.LoRAParam tensors, collect lora a/lora b pairs per module. 5. Load the raw base safetensors weights into a mutable NumPy dict. 6. Apply the LoRA delta: Wmerged = Wbase + αr · AB 7. Save the merged dict as a single model.safetensors file.
4.2
Tensor Shape Mapping
The delta computation must account for different tensor layouts between JAX (Tunix) and HuggingFace: Table 4: LoRA delta computation per module type. Module
JAX shapes
Delta & layout
A: (in, r), B: (r, heads, ∆ = (AB)T → (out, in) head dim) A: (in, r), B: (r, 2, Split B on dim-1; K=B[:, 0, :], kv einsum kv heads, hd) V=B[:, 1, :]; ∆ = (AB)T for each attn vec einsum A: (heads, hd, r), B: (r, Flatten A to (heads*hd, in) r); ∆ = (AB) → (in, heads*hd) gate/up/down proj A: (in, r), B: (r, out) ∆ = (AB)T → (out, in) q einsum
The HuggingFace convention stores weight matrices as (out, in) (row = output neuron), while the JAX einsum convention is often the transpose. All deltas are transposed to match HuggingFace layout before adding to the base weight.
4.3
Why Not Use PEFT Merge Directly?
The PEFT merge and unload() function expects a HuggingFace PreTrainedModel with PEFT adapters. The Orbax checkpoint contains raw JAX arrays organized in a tree matching Tunix’s flax NNX module hierarchy, with no correspondence to HuggingFace’s model class. There is no off-the-shelf bridge; the custom merger is necessary.
7
h2loop.ai — Technical Report
5
Training Results
5.1
Configuration
TPU vs GPU: Gemma 4 31B LoRA SFT
Both runs used identical hyperparameters: 1,244 optimizer steps, effective batch size 8, sequence length 4,096, 1 epoch over 10K CodeV-R1 samples, LoRA rank=64, α=64, AdamW with cosine LR schedule (peak 1e-4, 100 warmup steps, decay to 0). Table 5: Training performance comparison.
5.2
Metric
TPU v5p-8
GPU 2×H100
Winner
Wall-clock time Throughput (tokens/sec) Time per sample Final training loss
3.34 hr 763 1.21 s 0.072
5.39 hr 486 1.94 s 2.258
TPU (1.61× faster) TPU (1.57× faster) TPU (1.60× faster) TPU
Loss Curves
Figure 1: Training loss over the full epoch (x-axis normalized to training progress 0–1). Both curves follow similar convergence trajectories. The large absolute loss difference (0.072 vs 2.258 final) reflects the different loss computation: the TPU recipe uses assistant-only loss on ˜15% of tokens (the Verilog output tokens), while the GPU recipe computes loss over the full sequence including prompt tokens. When normalized to the same starting/ending range, the convergence speed is comparable.
8
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
Figure 2: Gradient norm over training (log scale). Both runs are stable after warmup. The GPU run shows sharper spikes (up to 106× at step 100) attributable to the larger effective batch aggregating more diverse gradients into a single update. The TPU gradient norms remain in the 1–10 range throughout.
5.3
Why is TPU Faster?
The TPU v5p-8 achieves 1.61× faster wall-clock time despite having only 4 chips vs 2 H100 GPUs. Several factors contribute: 1. HBM bandwidth: Each TPU v5p chip has 2.765 TB/s HBM bandwidth vs the H100’s 3.35 TB/s, but TPU has 4 chips vs 2 GPUs. Total aggregate bandwidth: 11.06 TB/s (TPU) vs 6.7 TB/s (GPU). 2. ICI interconnect: TPU chips communicate via the ICI mesh at 900 GB/s bidirectional bandwidth, significantly faster than the H100’s NVLink at ˜900 GB/s total (but NVLink is shared across more links). 3. XLA fusion: XLA compiles the forward+backward pass into a single fused kernel per layer, eliminating the kernel launch overhead present in PyTorch’s eager mode, even with FSDP. 4. No torch.compile overhead: The GPU run disables torch.compile (enforce eager=True) to avoid compatibility issues with the heterogeneous Gemma 4 attention layers, leaving PyTorch in eager mode throughout.
5.4
Training Cost Table 6: Training cost comparison (on-demand pricing, GCP us-central1).
Metric
TPU v5p-8
GPU a3-highgpu-2g
Winner
On-demand hourly rate Training time Total training cost Cost per 1M tokens trained
$16.80/hr 3.34 hr $56.11 $6.10
$22.12/hr 5.39 hr $119.23 $12.68
TPU (24% cheaper/hr) TPU (1.61× faster) TPU (2.12× cheaper) TPU (2.08× cheaper)
The 2.12× cost advantage compounds the lower hourly rate and faster execution time. On a longer training run (e.g., full CodeV-R1 at 87K samples), the cost difference would scale 9
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
proportionally.
6
Evaluation Results
6.1
Benchmark
Both fine-tuned models are evaluated on NVlabs/verilog-eval (spec-to-RTL task): 156 problems, 5 samples per problem at temperature=0.8, top p=0.95. Each generated Verilog module is compiled with iverilog and simulated against a reference testbench. Pass@k is computed using the standard unbiased estimator. Table 7: Evaluation results on verilog-eval spec-to-RTL.
6.2
Metric
TPU-trained
GPU-trained
pass@1 pass@5 Problems evaluated Samples per problem Checkpoint step Inference engine
0.6410 0.7949 156 5 1,244 (final) Tunix Sampler (JAX)
0.6974 0.8141 156 5 1,250 (final) vLLM (CUDA)
Per-Problem Analysis
Figure 3: pass@1 score per problem for both models. Most problems are solved at 100% by both. Differences concentrate on harder FSM, K-map, and cellular automata problems (indices 100–156).
10
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
Figure 4: Distribution of pass@1 scores across 156 problems. GPU-trained model: 101 problems at 100%, 39 at 0%. TPU-trained model: 89 at 100%, 50 at 0%. The difference is primarily in the zero-score bucket (11 additional problems failing entirely for TPU). The gap is within normal run-to-run variance for stochastic sampling at n=5 and is not statistically significant enough to attribute to hardware.
7
TPU Inference: vLLM on v6e-8
7.1
Setup Challenges
Serving Gemma 4 31B on TPU v6e-8 via vLLM required navigating several constraints that do not exist on the GPU path: 7.1.1
Docker-Only Deployment
The vLLM TPU build requires specific versions of JAX, libtpu, and HuggingFace libraries that conflict with pip-installed packages. The only working approach is the official Docker image (vllm/vllm-tpu:gemma4). Key flags: • --privileged: required for TPU device access from inside the container. • --network host: required for the gRPC TPU coordination. • --entrypoint vllm: the container’s default entrypoint is a wrapper script that parses args differently; using vllm directly is necessary. • --disable chunked mm input: Gemma 4 has heterogeneous attention head dimensions (head dim=256 for sliding, 512 for global) which breaks vLLM’s chunked multi-modal prefill path. This flag disables it. • Model specified as HuggingFace repo ID (google/gemma-4-31B-it), not a local path: vLLM-TPU rejects local paths due to HuggingFace Hub validation restrictions. 7.1.2
XLA Compilation on First Startup
vLLM-TPU pre-compiles XLA graphs for a set of padded input lengths (token buckets from 16 to 2048) at startup. This takes approximately 8 minutes on the first start. The compiled XLA cache is not persisted across container restarts by default, so every cold start incurs this cost. For production, mounting a persistent volume at the cache directory eliminates this overhead on warm restarts.
11
h2loop.ai — Technical Report
7.1.3
TPU vs GPU: Gemma 4 31B LoRA SFT
HBM Allocation
With tp=8 on v6e-8 (31.25 GB HBM per chip), the 31B model in bf16 needs approximately 31B × 2 bytes = 62 GB, spread across 8 chips = 7.75 GB/chip for weights. The remaining ˜23 GB/chip is available for KV cache. At 92% HBM utilization observed in benchmarks, this is a comfortable allocation.
7.2
GPU vLLM Setup (for comparison)
The GPU setup is more straightforward. It loads the model from a local path (the standard HuggingFace safetensors model pre-downloaded to ~ /models). With 2×80 GB = 160 GB HBM, the model weights use 62 GB, leaving 98 GB for KV cache in bfloat16. The GPU uses max model len=16384 and vLLM version 0.20.2 (vllm/vllm-openai:latest), compared to TPU’s version 0.19.0 (vllm/vllm-tpu:gemma4).
8
Inference Results and Analysis
8.1
Benchmark Configuration
The benchmark uses vllm bench serve with a random dataset, temperature=0, sweeping QPS from 1 to 64 across multiple input length regimes (512–15360 tokens input, 256–512 tokens output), 30 prompts per run. Both servers configured with max model len=16384, no chunked prefill. TPU uses TP=8; GPU uses TP=2. KV cache dtype: fp8 e5m2 (auto) on TPU, bfloat16 on GPU—giving TPU 2× the effective KV cache capacity. Table 8: Inference benchmark results across context lengths. QPS sweep peak = best sustained throughput across all QPS settings. Metric
TPU v6e-8
Short context (512 in / 256 out) — burst, QPS=inf Peak output throughput 1,403 tok/s Median TTFT 45 ms
GPU 2×H100
Winner
1,490 tok/s 51 ms
H100 (+6%) TPU (1.1×)
Medium context (1024 in / 512 out) — QPS sweep peak Peak output throughput 1,404 tok/s 1,387 tok/s Median TTFT @ QPS=4 49 ms 58 ms Saturates at QPS ∼32 ∼8
TPU (+1%) TPU (1.2×) TPU (4× more headroom)
Long context (4096 in / 512 out) — sustained load Peak output throughput 1,206 tok/s Median TTFT @ QPS=4 61 ms Median TPOT 23.9 ms
728 tok/s 1,443 ms 31.2 ms
TPU (+66%) TPU (23.6× faster) TPU (1.3×)
Very long context (8192 in / 512 out) Peak output throughput 482 tok/s Median TTFT @ QPS=4 1,013 ms Median TPOT 31 ms
449 tok/s 7,202 ms 45 ms
TPU (+7%) TPU (7.1× faster) TPU (1.5×)
Max context (∼16k in / 512 out) Peak output throughput Median TTFT @ QPS=4 Median TPOT
474 tok/s 62 ms 13.5 ms
326 tok/s 99 ms 19 ms
TPU (+45%) TPU (1.6×) TPU (1.4×)
vLLM version KV cache dtype
0.19.0 (vllm-tpu) fp8 e5m2 (auto)
0.20.2 (vllm-openai) bfloat16
— TPU (2× KV capacity)
12
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
Figure 5: Inference comparison across throughput, latency, and cost dimensions at the 128-token short-context baseline.
8.2
Explaining the Short-Context GPU Advantage
At short input lengths (≤2048 tokens), GPU edges ahead by ∼6% on peak output throughput. This is the memory-bandwidth-bound decode regime where the H100’s higher per-GPU HBM bandwidth (3.35 TB/s vs ∼1.6 TB/s per v6e chip) is most relevant. With short KV caches, each decode step reads relatively little KV state per token, so the bottleneck is weight loading—where H100 wins per chip.
8.3
Explaining the Long-Context TPU Dominance
At 4096+ token inputs, TPU v6e-8 comprehensively outperforms the GPU: 66% higher throughput and 23.6× faster TTFT at QPS=4. Two factors drive this: 1. Prefill compute: TTFT is dominated by the prefill phase (processing all input tokens), which is compute-bound —proportional to O(n2 ) for full-attention layers and O(n) for sliding-window layers. The v6e-8’s 8 chips deliver substantially more raw compute than 2 H100s for this workload. At 4096 input tokens, the GPU’s TTFT of 1,443 ms indicates it is fully saturated on compute during prefill; the TPU handles the same load in 61 ms. 2. fp8 KV cache: The TPU vLLM build uses fp8 e5m2 KV cache dtype by default, halving the memory footprint of each KV entry compared to the GPU’s bfloat16. With 250 GB total HBM on v6e-8 and half the KV storage cost, the TPU can hold 2× as many concurrent long-context requests in its KV cache before eviction. The GPU hits KV cache pressure at ∼8 concurrent long requests (QPS saturation at 8); the TPU sustains up to ∼32. The TensorCore utilization at 19.3% (measured by tpu-info at short-context decode) confirms the memory-bound regime for short sequences. At long contexts, utilization rises as prefill computation dominates.
8.4
Explaining the TTFT Advantage at Short Context
Even at short context (TTFT of 45 ms vs 51 ms), TPU is marginally faster. This comes from: 1. Higher aggregate compute: 8 v6e chips vs 2 H100s means more parallel matrix multiply units for the prefill pass.
13
h2loop.ai — Technical Report
TPU vs GPU: Gemma 4 31B LoRA SFT
2. XLA fusion: XLA fuses QKV projection with the sliding-window and full-attention kernels in Gemma 4’s heterogeneous attention, reducing HBM roundtrips. The GPU uses TRITON ATTN backend due to the mixed head-dim layout, which is not as aggressively fused.
8.5
Explaining Slightly Better GPU P99 Tail at Short Context
At 128-token inputs, GPU achieves P99 ITL of 20.49 ms vs 22.29 ms on TPU. This is due to vLLM-TPU’s XLA bucket padding: inputs are padded to the nearest power-of-2 token bucket (16, 32, 64, 128, 256, ...). Requests that land near a bucket boundary receive unnecessary padding tokens, inflating their per-step compute and causing P99 spikes. The GPU’s CUDA path does not have this constraint.
8.6
Inference Cost Table 9: Inference cost comparison by context length (on-demand).
Workload
TPU tok/hr
GPU tok/hr
TPU $/1M tok
GPU $/1M tok
Short context (≤2048 tokens) Long context (4096 tokens)
5.05M 4.34M
5.36M 2.62M
$4.27 $4.95
$4.13 $8.44
Hourly rate
$21.52/hr
$22.12/hr
For short-context workloads, GPU is 3% cheaper per token. For long-context workloads (≥4096 tokens), TPU is 41% cheaper per output token due to the combined effect of higher throughput and comparable hourly rate.
9
Total Cost of Ownership Table 10: End-to-end cost comparison (on-demand, GCP us-central1/asia-northeast1). Cost Component
TPU
GPU
Training (v5p-8 / a3-highgpu-2g) Inference — 1 hr serving (v6e-8 / a3-highgpu-2g) Total (train + 1 hr inference)
$56.11 $21.52 $77.63
$119.23 $22.12 $141.35
Train + 8 hr inference Train + 24 hr inference
$228.27 $572.59
$296.19 $650.11
TPU is 1.82× cheaper end-to-end for a single training run plus a day of inference serving at short context, with the advantage driven primarily by training cost savings. At long-context workloads (≥4096 tokens), where TPU is 41% cheaper per output token, the inference savings compound further.
14
h2loop.ai — Technical Report
10
TPU vs GPU: Gemma 4 31B LoRA SFT
Summary
Figure 6: TPU vs GPU relative advantage across all dimensions (ratio > 1.0 = TPU better). Training cost and TTFT are the largest wins; eval quality is the only metric where GPU edges ahead.
10.1
TPU Advantages
1. Training speed: 1.61× faster wall-clock time (3.34 hr vs 5.39 hr) 2. Training cost: 2.12× cheaper ($56.11 vs $119.23) — lower hourly rate compounded by faster execution 3. Inference TTFT at 4096 tokens: 23.6× faster (61 ms vs 1,443 ms at QPS=4) 4. Inference TTFT at 8192 tokens: 7.1× faster (1,013 ms vs 7,202 ms at QPS=4) 5. Inference throughput at long context (4096 tokens): +66% output tok/s (1,206 vs 728) 6. Inference throughput at max context (∼16k tokens): +45% output tok/s (474 vs 326) 7. Memory headroom: 411 GB total HBM (training) + fp8 KV cache enables 2× more concurrent long-context requests
10.2
GPU Advantages
1. Ecosystem maturity: PyTorch + HuggingFace has broader library support, simpler LoRA injection, and faster prototyping with fewer workarounds 2. Short-context throughput: marginally higher peak burst at ≤2048 token inputs (H100 +6%) 3. Short-context cost: 3% cheaper per token at short context due to higher throughput 4. Eval quality (marginal): +5.6pp on pass@1, attributable to training objective differences rather than hardware 5. P99 ITL at short context: slightly better tail latency (20.49 ms vs 22.29 ms at 128-token inputs) due to no XLA bucket padding overhead
15
h2loop.ai — Technical Report
10.3
TPU vs GPU: Gemma 4 31B LoRA SFT
Conclusion
Training: TPU is the clear winner — 1.61× faster and 2.12× cheaper. The engineering overhead of porting from PyTorch to JAX + Tunix is real (roughly a week for mesh configuration, sharding annotation fixes, custom data pipeline, and checkpoint conversion) but is a one-time cost amortized across many runs. Inference: TPU is the clear winner at long-context workloads. At 4096-token inputs, TPU delivers 66% higher throughput and 23.6× faster TTFT. At short context (≤2048 tokens), performance is comparable with a marginal GPU throughput edge (+6%). The fp8 KV cache on TPU doubles effective KV capacity, enabling 4× higher QPS saturation at medium context. The total cost advantage of 1.82× over a representative train+serve pipeline makes TPU the preferred infrastructure choice for production ML workloads at h2loop.ai’s scale.
10.4
Resources
• TPU — Recipe & Eval Scripts: https://github.com/h2loop/gemma-tpu • Eval benchmark: https://github.com/NVlabs/verilog-eval
References [1] Y. Zhu et al., “CodeV: Empowering LLMs for Verilog Generation through Multi-Level Summarization,” arXiv:2407.10424, 2024. [2] M. Liu et al., “VerilogEval: Evaluating LLMs for Verilog Code Generation,” ICCAD, 2023. [3] Google DeepMind, “Gemma 4 Technical Report,” 2025. [4] E. Hu et al., “LoRA: Low-Rank Adaptation of Large Language Models,” ICLR, 2022. [5] W. Kwon et al., “Efficient Memory Management for Large Language Model Serving with PagedAttention,” SOSP, 2023.
16