Gemma 4 QAT on One TPU v5e: What Runs and What Doesn't
This article provides a step by step guide to repacking Google's quantization-aware-trained (QAT) Gemma 4 weights for vLLM and serving them on one Google Cloud TPU v5e chip, with every build scored for classification, math, tool calling, throughput and long prompts. Every per-record output, log and script is committed. On one v5e chip the repacked QAT builds serve every Gemma 4 size from E2B to 26B. Through 12B they read level with bf16 on a 3,880-record classification suite, they score up to 2.4 points above Google's own 4-bit exports at the same speed, and they make 12B the largest model on the chip: 11.31 GiB of weights, 675 output tokens per second at 16 requests, 0.964 on GSM8K and 0.955 on BFCL tool calling. https://github.com/xbill9/gemma4-dev/tree/main/jev-tpu-v5e1
Why Repack?
One TPU v5e chip ( v5litepod-1 ) has 15.75 GiB of HBM. Gemma 4 E4B at bf16 is 14.9 GiB and 12B is 22.4 GiB, so everything above E2B needs 4-bit or 8-bit weights on this chip. Google trained 4-bit versions of every Gemma 4 size and publishes the trained values as "unquantized" bf16 checkpoints ( -qat-q4_0-unquantized ). Every group of 32 weights in them already sits on a 16‑level grid. A repack stores those values in a format vLLM serves, without re‑rounding them:
| Build | What is stored |
|---|---|
| q4w4a16 | int4 weights holding the QAT grid exactly, 16‑bit activations |
| q4w4a16emb4 | the same, with the vocabulary tables ( embed_tokens , lm_head , per‑layer embeddings) also int4 |
| w8a8 | int8 weights per channel from the QAT values, int8 activations per token |
| w8a8emb4 | the same, with int4 vocabulary tables |
v5e multiplies int8 by int8 natively, which makes the int8 builds the fast ones on this chip.
At This Point You Should Have…
- A Google Cloud project with TPU v5e flex‑start quota in us‑west4‑a, and the gcloud CLI logged in
- A Cloud Storage bucket for checkpoints and results
- A Hugging Face token in Secret Manager as hf‑token
- A clone:
git clone https://github.com/xbill9/gemma4-dev
Step 1 - Repack the QAT Weights
The 4‑bit repack recovers each group's trained step and stores the group as int4, so every value keeps its trained place on the grid:
huggingface-cli download google/gemma-4-12B-it-qat-q4_0-unquantized \
--local-dir ~/models/gemma-4-12B-it-qat-q4_0-unquantized
python3 ../jev-tpu-31b/repack_q4_0.py repack ~/models/gemma-4-12B-it-qat-q4_0-unquantized \
~/models/gemma-4-12B-it-qat-q4_0-w4a16-ct
The int8 builds take the same QAT values to int8 per channel with ../jev-tpu-31b/w8a8_from_qat.py. The published builds are on Hugging Face under xbill9/gemma-4-*-it-qat-*. Serving them on the TPU backend uses three additions to vLLM's tpu_inference, in jev-tpu-v5e1/patches/:
- an int8
W8A8method on the JAX path, - int4 embedding tables that stay packed on the chip and unpack only the rows a step reads,
- an int4
lm_head.
Step 2 - Serve on One v5e Chip
Each run is a flex‑start queued resource that boots, applies the patches to the pinned vLLM image, serves a list of builds one after another and uploads every result:
gcloud alpha compute tpus queued-resources create jev-tpu-v5e1-$RUN \
--zone us-west4-a \
--accelerator-type v5litepod-1 --runtime-version v2-alpha-tpuv5-lite \
--node-id jev-tpu-v5e1-$RUN-node \
--provisioning-model flex-start --max-run-duration 4h --valid-until-duration 2h \
--metadata jev-code=<bundle
Comments
No comments yet. Start the discussion.