{"slug": "gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact", "title": "Gemma 4 E2B in Pure JAX on a Colab TPU: Google's 4-Bit Export Against an Exact Repack", "summary": "A developer published an Apache 2.0 Colab notebook that serves Gemma 4 E2B on a single TPU v5e chip using a pure-JAX engine and compares Google's 4-bit QAT export against an exact repack of the trained quantization grid. The repack lands 342.6 times closer to the original model in next-token predictions, matches its top token 99.29% of the time versus 85.96%, and downloads 0.8 GB less because Google's export redundantly stores lm_head.weight as a copy of the token embedding.", "body_md": "This article provides a step by step guide to a Colab notebook that serves Gemma 4 E2B on a single TPU v5e chip with a pure-JAX engine and compares two 4-bit builds of the same model against the weights Google trained. Every number below was measured in the notebook on a Colab v5e-1 runtime, and the executed notebook is committed.\n\nGoogle ships E2B in 4 bits as `gemma-4-E2B-it-qat-w4a16-ct`. Its export rounds every weight a second time, onto a grid the model never trained on. A repack that stores the trained grid instead lands 342.6 times closer to the original model in next-token predictions, matches its top token 99.29% of the time against 85.96%, writes the same 128-token story word for word, runs at the same speed and downloads 0.8 GB less.\n\nThis notebook is an entry in the AI GDE Marathon: JAX on TPU Tutorial, a series of Apache 2.0 Colab notebooks covering JAX on the TPU backend. Colab's TPU runtimes are the same v5e and v6e silicon the source measurements came from, so every cell measures on the reader's chip and prints the source result beside it as context.\n\nGemma 4 E2B was trained with quantization-aware training (QAT): during training every weight was held on a 4-bit grid, one scale per group of 32 values, so the model learned to work with exactly those values. Google publishes the result twice:\n\n| Build | Hugging Face | What it stores | \n|---|---|---|\n| QAT reference | `google/gemma-4-E2B-it-qat-q4_0-unquantized` | the trained grid values, in bf16 | \n| Google's export | `google/gemma-4-E2B-it-qat-w4a16-ct` | int4, each group re-rounded with step = largest weight ÷ 7.5 | \n| The repack | `xbill9/gemma-4-E2B-it-qat-q4_0-w4a16-ct` | int4, the trained levels and the trained step | \n\nBoth 4-bit builds use the same compressed-tensors W4A16 format and the same shapes, so any loader that reads one reads the other. The difference is which numbers are inside.\n\nNo Hugging Face token is needed: all three checkpoints are public and ungated, and the engine is a public GitHub repo.\n\nOpen the notebook from the link above, then Runtime → Change runtime type → v5e-1 TPU, then Runtime → Run all. Cell 1.1 records the chip:\n\n```\nJAX 0.7.2\n1 device(s): TPU v5 lite (tpu)\nHBM 16.91 GB   host RAM 50.5 GB   free disk 195 GB\n```\n\nCell 1.2 clones the engine, [`tpu-jax`](https://github.com/xbill9/tpu-jax), at a pinned commit. It is a Gemma 4 E2B decoder in pure JAX with no PyTorch in the path, and its 4-bit path unpacks each weight with plain XLA operations before multiplying.\n\nEach checkpoint is pinned to a Hub revision, so the bytes you download are the bytes measured here. Cell 2.2 then indexes every tensor:\n\n```\nqat      1951 tensors     0 packed int4   lm_head stored: False\nstock    2504 tensors   276 packed int4   lm_head stored: True\nrepack   2503 tensors   276 packed int4   lm_head stored: False\n\nstock lm_head identical to embed_tokens: True   (0.805 GB)\nstock minus repack on disk: 0.805 GB\n```\n\nBoth 4-bit builds pack the same 276 linear layers. Google's export also stores `lm_head.weight`, a byte-for-byte copy of the token embedding, although the config ties the two. That one tensor is the whole 0.8 GB difference in download size.\n\nCell 3.1 reads 24 tensors from the first, middle and last layers, every projection kind and 5.7 million groups of 32, from all three files. It unpacks each 4-bit build with the engine's own function and compares the result with the QAT values:\n\n```\nall 24 tensors\n  stock   groups 5,738,496   rel err 0.0667   identical 12.08%   step=max/7.5 100.00%   peak on a level 0.00%\n  repack  groups 5,738,496   rel err 0.0019   identical 73.65%   step=max/7.5 0.00%   peak on a level 99.92%\n\n  repack: level the largest weight of each group sits on, share of groups\n   1: 0.0%   2: 0.0%   3: 0.0%   4: 0.0%   5: 0.0%   6: 0.0%   7: 39.5%   8: 60.5%\n```\n\nGoogle's export uses largest weight ÷ 7.5 as the step in every group, and its values sit 6.7% from the trained ones. The repack sits 0.19% away, three values in four are bit-for-bit the trained value, and the rest differ only by the bf16 rounding of the stored step.\n\nThe last line explains why the trained step has to be stored. If a group's largest weight sits on level *m* of the trained step *d*, the ÷ 7.5 rule gives a step of *m·d* / 7.5, which never equals *d* because *m* is a whole number from 1 to 8. The largest weight sits on level 7 in 39.5% of groups and level 8 in the other 60.5%, so a fixed ÷ 8 rule would be right for only some of them.\n\nCells 4.2 to 4.4 load each model in turn, feed all three the same 8 WikiText-2 sequences of 512 tokens, and keep each model's predicted next-token distribution at every position. Cell 4.5 compares each 4-bit build with the QAT model on the same 4,096 positions:\n\n```\nbuild     mean KL  same top token  perplexity ratio   per-sequence mean KL, min to max\nstock     0.07141          85.96%            1.0643   0.05807 to 0.08348\nrepack    0.00021          99.29%            0.9994   0.00015 to 0.00028\n\npositions where the repack is closer to the QAT model than the stock build: 100.0% of 4,096\nstock KL / repack KL: 342.6x\n```\n\nThe repack is closer to the QAT model at every one of the 4,096 positions, and the gap holds on every sequence: the stock build's best sequence is still 200 times further off than the repack's worst. Google's export also predicts the text 6.4% worse than the model it was made from.\n\nCell 5.1 decodes a fixed prompt, \"The history of the Roman Empire\", for 128 greedy tokens, once to compile and three times timed:\n\n```\nbuild    disk GB  weights GB  decode tok/s  tokens equal to QAT greedy\nqat        10.21        9.26         135.0                  128 of 128\nstock       8.32        6.56          93.7                    5 of 128\nrepack      7.51        6.56          93.3                  128 of 128\n\nrepack / stock decode speed: 0.996x\n```\n\nThe two 4-bit builds tie on speed and put the same 6.56 GB on the chip, because the engine never loads the duplicate `lm_head`. The repack writes the QAT model's story token for token. Google's export opens with the same five words, \" is a vast and complex\", then writes \"narrative\" where the QAT model writes \"tapestry\", and the two stories diverge from there.\n\nThe bf16 QAT model decodes faster than either 4-bit build in this engine, because the engine unpacks every 4-bit weight at every step. The engine can unpack once at load instead (`dequant_at_load=True`), which trades the HBM saving for that speed.\n\nCell 6.1 prints the earlier vLLM measurements beside the ones from the notebook:\n\n```\n                                           source (vLLM)               today (pure JAX)\n                                            stock         repack          stock         repack\nrelative error vs QAT values        0.0665-0.0667            n/a         0.0667         0.0019\ngroups with step = max/7.5                100.00%            n/a        100.00%          0.00%\nmean KL vs QAT model                          n/a            n/a        0.07141        0.00021\nsame top token as QAT model                   n/a            n/a         85.96%         99.29%\ntest suite, 3,880 records                   65.5%          67.8%        not run        not run\ndownload, GB                                 8.32           7.51           8.32           7.51\ndecode tok/s, 1 request                     136.6          136.5           93.7           93.3\ndecode speed, repack / stock                              0.999x                        0.996x\n```\n\nThe absolute speeds differ between the pairs because the engines differ: vLLM's int4 kernel against a reference unpack-then-multiply in pure JAX. The ratio inside each pair is what carries, and it is 0.999x under vLLM and 0.996x here. The same holds for quality: under vLLM on a 3,880-record test suite the repack scored 67.8% against 65.5%, 2.4 points higher (95% range +1.4 to +3.4) and within 0.6 points of the bf16 release.\n\n|  | Google's export | The repack | \n|---|---|---|\n| Error against the trained weights | 🔴 6.7% | 🟢 0.19% | \n| Mean KL against the QAT model | 🔴 0.07141 | 🟢 0.00021 | \n| Same top token as the QAT model | 85.96% | 🟢 99.29% | \n| Greedy story tokens equal to QAT | 🔴 5 of 128 | 🟢 128 of 128 | \n| Download | 8.32 GB | 🟢 7.51 GB | \n| Weights on the chip | 6.56 GB | 6.56 GB | \n| Decode tok/s, one request | 93.7 | 93.3 | \n| Format | compressed-tensors W4A16 | compressed-tensors W4A16 | \n\nThe repack. It holds the weights Google trained, runs at the same speed in the same format, puts the same bytes on the chip and is 0.8 GB smaller to download. Any loader that reads Google's `-qat-w4a16-ct` reads it unchanged, so switching is a change of repo name.\n\nThe goal of this notebook was to measure, on a reader's own Colab TPU, whether Google's 4-bit Gemma 4 E2B export holds the weights the model was trained on, and what a repack that holds them exactly changes. The key to the solution was putting all three checkpoints through one pure-JAX loader, one forward pass, one text and one prompt, with the QAT checkpoint as the reference for both 4-bit builds. The results were:\n\nScope: one Colab v5e-1 runtime (one TPU v5 lite chip, 16.91 GB HBM), JAX 0.7.2, `tpu-jax` at commit `4b9f8e9`, run on 2026-10-09. The output comparison uses 8 WikiText-2 test sequences of 512 tokens, and speed is the median of three 128-token greedy runs of one prompt. The 3,880-record test suite and the vLLM speeds are earlier results on a v5e chip, cited as context. The repack is unofficial and derived from Google's release under Apache 2.0. Parts of the analysis and writing were done with AI assistance (Claude); every figure comes from the committed executed notebook.\n\nThe comparison of Google's 4-bit Gemma 4 E2B export against an exact repack was validated on a Colab TPU with an incremental step by step approach.", "url": "https://wpnews.pro/news/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact", "canonical_source": "https://dev.to/gde/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-googles-4-bit-export-against-an-exact-repack-4dle", "published_at": "2026-10-10 00:51:07+00:00", "updated_at": "2026-10-10 00:58:02.585779+00:00", "lang": "en", "topics": ["large-language-models", "machine-learning", "ai-infrastructure", "ai-tools", "developer-tools"], "entities": ["Google", "Gemma 4 E2B", "JAX", "Hugging Face", "TPU v5e", "xbill9/tpu-jax", "xbill9/gemma-4-E2B-it-qat-q4_0-w4a16-ct", "google/gemma-4-E2B-it-qat-w4a16-ct"], "also_reported_by": [], "alternates": {"html": "https://wpnews.pro/news/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact", "markdown": "https://wpnews.pro/news/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact.md", "text": "https://wpnews.pro/news/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact.txt", "jsonld": "https://wpnews.pro/news/gemma-4-e2b-in-pure-jax-on-a-colab-tpu-google-s-4-bit-export-against-an-exact.jsonld"}}