Gemma 4 QAT on One TPU v5e: What Runs and What Doesn't
DEV Community

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 W8A8 method 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
Read on DEV Community ↗ ← Back to News

Comments

No comments yet. Start the discussion.