cd /news/machine-learning/gemma-4-in-pure-jax-what-changes-bet… Β· home β€Ί topics β€Ί machine-learning β€Ί article
[ARTICLE Β· art-116105] src=dev.to β†— pub= topic=machine-learning verified=true sentiment=Β· neutral

Gemma 4 in Pure JAX: What Changes Between Turing and Ada, and What Doesn't

A developer's hand-written Gemma 4 port in pure JAX runs on both Turing and Ada NVIDIA GPUs with identical weights, but performance differs drastically due to hardware-specific compute dtype handling and memory constraints. The port avoids Triton's shared-memory ceiling by using XLA for attention, yet a wrong compute dtype on pre-Ampere GPUs silently emulates through fp32, costing 87% of decode speed. The developer's solution reads the device's compute capability to select the correct dtype, ensuring the fast path runs only on Ada.

read8 min views1 publishedAug 31, 2026

This article is a measurement report on running a hand-written Gemma 4 port in

pure JAX across two NVIDIA GPUs a generation apart, and on the two places the

"it's just JAX" abstraction leaks. One of those leaks costs 87% of decode and

nothing in the logs is red.

The code is here:

https://github.com/xbill9/gemma4-dev

One port, one build, one checkpoint, two cards. Everything below comes from two

archived runs, named so you can check them.

G5g G6
Chip NVIDIA T4G β€” Turing, SM 7.5, 15,360 MiB
NVIDIA L4 β€” Ada, SM 8.9, 23,034 MiB
Host
g5g.2xlarge spot β€” Graviton2, aarch64
g6.2xlarge spot β€” x86_64, us-east-1d
Checkpoint
google/gemma-4-E2B-it , dense reference
google/gemma-4-E2B-it , dense reference
Compute dtype
float16 (device-chosen)
bfloat16 (device-chosen)
Stack jax 0.11.1, CUDA from pip jax 0.11.1, Python 3.14
Run cited 2026-08-28-full-run-cached-g5g
2026-08-28-first-serve-g6

Build id 51bc52c9e2e9

on both, config ple4 + int8_lm_head

, and

tpu_jax_weight_bytes

reads 6,155,450,950 on both cards β€” the same integer.

Only the chip and its host differ.

g5g.2xlarge

and g6.2xlarge

google/gemma-4-E2B-it

The port lives in ports/gemma4/

and is driven by a generation loop behind an

OpenAI-compatible server. No PyTorch, no vLLM, no torch_xla

.

The premise under test is that the same source runs on both cards with nothing

changed but a config file. It mostly holds. The interesting part is where it does

not.

Any port has to carry four irregularities, and none of them are optional.

head_dim=256

, global layers use That first irregularity is the expensive one. On the vLLM path the heterogeneous

head dims force the Triton attention backend:

Gemma4 model has heterogeneous head dimensions
(sliding=256, global=512); falling back to the Triton attention backend

On a Turing GPU that backend then asks for shared memory the hardware does not

have:

triton.runtime.errors.OutOfResources: out of resource: shared memory,
Required: 147456, Hardware limit: 65536

JAX never enters that conversation. Attention is ordinary XLA rather than a

hand-tiled kernel, so there is no per-block shared-memory ceiling in the attention

path at all. The irregular geometry that is a special case everywhere else is just

array shapes here.

This is the single most expensive lesson in the repository.

A wrong compute dtype does not raise. It emulates. bfloat16

on a pre-Ampere

GPU does not fail β€” XLA routes it through fp32 and most of decode disappears into

conversion. Nothing in the logs is red.

So the port does not take the dtype from a config file. It reads the live compute

capability off the device:

COMPUTE_DTYPE = float16 if IS_PRE_AMPERE else bfloat16

On the SM 8.9 Ada card that resolves to bfloat16

. On the SM 7.5 Turing card it

resolves to float16

β€” Turing's only real 16-bit datapath, since it has neither

bf16 nor fp8.

The first line the server emits is the policy, so a misconfiguration is one grep

away rather than a mystery in the throughput:

INFO ports.gemma4.jax_e_model: jax_e_model device policy: platform=gpu
compute_capability=8.9 compute_dtype=bfloat16 pallas_interpret=False

pallas_interpret=False

matters just as much. It is the difference between

serving and silently running a simulator.

Here is the part that does not port, and it is not a bug. It is a real hardware

difference wearing a portable API.

The fused W4A16 kernel is written in Pallas, and it was tiled for a device with

16 MB of scratchpad per core. At this model's shapes the tiles want 550 KiB to 1.1 MiB per block.

On a GPU, Pallas lowers through Triton, and those tiles become shared memory.

Turing gives you 64 KiB per block. Ada raises the ceiling, but nowhere near a

megabyte.

So the fast path runs on neither card. The engine computes the requirement at

startup and refuses with the arithmetic attached, rather than dying as a cryptic

OutOfResources

at the first token:

check_w4a16_fits_scoped_memory()

The practical consequence is that both GPU rigs serve the dense reference checkpoint at 16-bit.

A padding-eviction bug in the KV ring cache cost a week, and it is the kind only

Gemma 4's geometry produces.

The invariant is that a cache index is an absolute real position, and padding never occupies an index a real position uses. A port that right-pads into the

200

, status: "success"

, and output likeThe The The The

.Nothing in the logs is red. Nothing in the metrics is red. The only thing that

catches it is a degeneracy check on the output itself, which the server now runs on

every response.

The scariest bugs in this project all returned success.

jax[cuda13]

supplies CUDA as wheels, so the install needs no CUDA toolkit, no

Rust, and no compiler on the box.

Install: 117 s, with the cache restore included

XLA's persistent compilation cache ports as-is. On the T4G rig it restores 805 files / 12 MB in 6 seconds onto a fresh instance, from a box that had already

max_new_tokens

is a static_argnames

entry, so (bucket, max_tokens)

is the

compiled shape on every backend. A harness that does not warm up misreports the

rig badly.

On the T4G the first request off a fresh engine took 18.06 s against 4.50 s warm β€” a 4.0x whole-request ratio, measured in

2026-08-21-cuda13-py314-g5g

.That run also notes something worth repeating: the 56x figure from the first-serve

baseline is TTFT specifically, not the same measurement as the whole-request

ratio. They are not interchangeable.

64 output tokens, concurrency 1, 3 repeats per cell, median. "Decode, gauge" is the

engine's steady-state counter. "End-to-end" is wall time over the whole request,

prefill included.

Input tokens T4G gauge T4G end-to-end L4 gauge L4 end-to-end
41 12.9 tok/s 12.43 tok/s 48.5 tok/s
46.23 tok/s
521 13.0 tok/s 11.28 tok/s 48.4 tok/s
42.87 tok/s
2,057 12.9 tok/s 8.22 tok/s 48.3 tok/s
34.57 tok/s
3,593 β€” β€” 48.3 tok/s
27.55 tok/s

Decode moves 0.8% across a 50x context range on the T4G and 0.4% on the L4. End-to-end

falls hard on both.

That fall is prefill being linear in the padded bucket, not decode degrading. They

are two different claims, and conflating them makes a benchmark a lie. Quote the gauge.

A cost proportional to the weights rather than the context produces exactly this

shape, which is why KV is not what sets decode speed on either card β€” despite

Gemma 4's whole KV story.

On context specifically: MAX_MODEL_LEN=4096

is the honest number on the T4G.

4,105 prompt tokens serve; 5,120 fails on a prefill transient.

Profiling decode with xprof on the Turing card, 20 decode steps with the service

stopped:

πŸ₯‰ T4G (SM 7.5) πŸ₯‡ L4 (SM 8.9)
dtype conversion 54.1% 0.0%
fp32 gemvx
32.8% absent
Tensor Core 0.0% 0.0%
Total kernel time 1,466.0 ms 362.8 ms
Decode, gauge 12.9 tok/s 48.4 tok/s
Peak HBM bandwidth 298.083 GiB/s 279.441 GiB/s
Share of bandwidth roofline 26% ~100%

1,466 ms of kernels across 108 distinct kernels on a Tensor Core GPU, without one

Tensor Core firing. More than half of decode went to converting numbers between

formats before any math happened.

The obvious hypothesis was bf16 weights being converted on a chip with no bf16

datapath. So the checkpoint was converted to float16 host-side and re-run.

Parameter dtypes read {'float16': 541, 'uint8': 1, 'int8': 1}

β€” and conversion stayed at 54.0%.

The measurement itself is solid. The same profile on a different instance, a

different AMI and a restored cache landed at 1466.0 ms against 1467.1 ms. 1.1 ms apart on 1467.

The Ada card resolves it. Converting the stored weights changed nothing because

storage dtype was never the problem: Turing has no native bf16, and the fp32

gemvx

line is the tell β€” XLA was round-tripping through fp32 regardless of what

the file on disk said.

Give it a card where storage and compute dtype actually match, and the 54%

conversion and the 32.8% fp32 path vanish together. An 87% tax gone, for 3.7x

the throughput, and a rig sitting at its bandwidth roofline instead of 26% of it.

The /health

endpoint on the L4 reports weights=bfloat16 activations=bfloat16

β€” storage dtype and compute dtype matching for

kv_cache=bfloat16 pre_ampere=false

the first time on this engine.

Tensor Core utilization is 0.0% on the Ada card too β€” 100 distinct kernels,

362.8 ms of them, and not one Tensor Core firing.

Removing the dtype pressure made the machine roughly four times faster without

making it touch the hardware it was sold for. That is the open question now, and it

is a better one than the question this started with.

Both rigs run on spot capacity and are terminated after collection. The XLA cache

is pushed to S3 before teardown, which is what makes the 6-second restore on a

fresh instance possible.

The goal of this article was to find out which parts of "it's just JAX" survive a

move between GPU generations. The key to the solution was reading the compute dtype

off the live device rather than a config file. The measured results were:

Scope: two spot instances, one in us-east-1d

, each measured once with 3 repeats

per sweep cell and medians reported. The two boxes differ in host architecture

(aarch64 against x86_64) and base image as well as in GPU, so this is not a

single-variable experiment; the payload is byte-identical across them β€” same build

51bc52c9e2e9

, same config, same 6,155,450,950 bytes of weights β€” which is the

basis for attributing the difference to the chip. The Turing profile was reproduced

on a second instance at 1466.0 ms against 1467.1 ms; the Ada profile was measured

once.

The strategy for using MCP for Gemma 4 serving across GPU generations was validated

with an incremental step by step approach.

── more in #machine-learning 4 stories Β· sorted by recency
── more on @gemma 4 3 stories trending now
sponsored brought to you by zahid.host 4,200+ EU-deployed projects
reading about agents? ship yours in a single git push.

Run your AI side-project on zahid.host

EU-based hosting, git-push deploys, automatic HTTPS, no cold starts. Free tier with a custom domain β€” perfect for shipping the agent you just read about.

$git push zahid main
β†’ Live at https://your-agent.zahid.host βœ“
Get free account β†’ Pricing
from €0/mo Β· no card required
LIVE [news/gemma-4-in-pure-jax-…] indexed:0 read:8min 2026-08-31 Β· β€”