Observed Signal · Aug 29, 2026 · Technical Release · Source: DEV Community · Impact: 3/5 · Sentiment: Positive

Gemma 4 in Pure JAX: TPU-to-GPU Port Lessons

Executive Signal Summary

A developer ported a Gemma 4 E2B checkpoint to a single pure-JAX codebase and ran it on Cloud TPU v5e/v6e and an NVIDIA T4G (on an AWS Graviton2 host). Most model code, compilation cache behavior, and static-shape discipline transferred unchanged, but two hardware-dependent issues surfaced: a fused W4A16 kernel written in Pallas (tiled for TPU VMEM) cannot run on GPUs due to much smaller shared-memory limits, and compute-dtype selection must be detected at runtime (float16 vs bfloat16) to avoid hidden conversion costs. The article documents Gemma 4’s four model irregularities, a KV-ring padding bug that produced silent token loops, measured decode throughput (13.10 tok/s on T4G), and profiling that shows unexpected conversion overhead on a Turing GPU.

Polaris7 AgentPolaris7 Strategic Assessment
High Confidence

Provides actionable, measured insights about LLM inference portability across TPU and GPU hardware (kernel memory-model mismatch, runtime dtype selection, and performance profiles) that matter to teams deploying foundation-model inference.

SIGNAL RADAR

Track NVIDIA Signals & Market Shifts in Real-Time

Polaris7 autonomous intelligence agents track regulatory filings, primary sources, executive changes, and deal flow 24/7. Create your free Explorer workspace to monitor these entities.

Start Free in Explorer
Free Explorer tierNo credit card requiredInstant watchlist setup

Key Takeaways & Evidence Grounding

  • A single pure-JAX Gemma 4 port was run on Cloud TPU v5e, Cloud TPU v6e, and an NVIDIA T4G attached to an AWS Graviton2 host.
  • Gemma 4 E2B requires four nonstandard accommodations: two attention head dimensions (256 and 512), 8:1 MQA, a KV-share map collapsing 35 layers onto 15 caches, and a 512-slot sliding ring with a 4.70 GB per-layer embedding table.
  • A fused W4A16 kernel implemented in Pallas is tiled for TPU VMEM (16 MB/core) and cannot execute on GPUs because GPU shared-memory per-block is orders of magnitude smaller.
  • The port must detect device compute capability at runtime and choose compute dtype (float16 vs bfloat16); on a Turing GPU, dtype conversion dominated runtime (~54% of decode time).
  • Measured decode throughput: 13.10 tokens/sec on the NVIDIA T4G; weights resident ~6.155 GB.

Connected Companies & Entities

5 Entities mapped

“The same code runs on Cloud TPU v5e and v6e, and on an NVIDIA T4G attached to an AWS Graviton2 host....”

“Model | `google/gemma-4-E2B-it`, dense reference build...”

“The port ... is driven by a generation loop behind an OpenAI-compatible server....”

“The code is here: [github.com/xbill9/gemma4-dev](https://github.com/xbill9/gemma4-dev)...”

Primary Source Grounding & Direct Attribution
Direct Origin Attribution
Primary Reporting: DEV Community•Published: Aug 29, 2026
Original Coverage Title: “Gemma 4 in Pure JAX: What Ports from TPU to GPU, and What Doesn't”

Related Market Signals & Shifts

Recent verified developments and strategic activity across this market segment.

Large Language Models (LLM) & AIJul 29, 2026

One TPU Chip, Eight Agents: Serving Small Agent Workloads

An engineer implemented a pure-JAX serving path to run a Gemma 4 E2B quantization-aware-trained (QAT) checkpoint on a single Cloud TPU v6e chip because vLLM could not load the QAT export on TPU. The author created a safetensors→JAX loader and a JAX decode kernel, validated correctness against full re-forward, and measured kernel decode rates up to ~2,888 tok/s (int4/int8 donated path). Memory math shows eight 8K contexts fit comfortably on a 32 GB HBM v6e chip (≈1.21 GB KV for eight 8K contexts). However, the experimental server lacks prefix caching, guided/schema-constrained decoding, and continuous batching, so end-to-end HTTP serving without batching reached only ~139–143 aggregate tok/s with latency rising under contention. Verdict: viable experimental path for cases that need the QAT checkpoint, but not yet a drop-in vLLM production replacement.

Read assessment
Large Language Models (LLM) & AIAug 29, 2026

Serving Google's Gemma 4 on AWS G5g with Pure JAX

A technical deployment guide demonstrating how to serve Google's open model Gemma 4 (google/gemma-4-E2B-it) on an AWS EC2 G5g instance (Graviton2 + NVIDIA T4G) using a pure JAX stack. The article includes a reproducible repo, prerequisites, automated install and verification steps, performance measurements (approximately 13 tokens/sec decode on the T4G with pure JAX versus 43 tok/s for a patched vLLM), and operational notes on AMI selection, Secrets Manager use, SSM Run Command administration, and S3-based XLA cache. The author emphasizes the 117-second reproducible deployment time and trade-offs between ease-of-deploy and raw throughput.

Read assessment
Large Language Models (LLM) & AIMay 24, 2026

Run Gemma 4 26B on GTX 1080 with llama.cpp

A developer how-to demonstrating how to run Google’s Gemma 4 26B‑A4B Mixture‑of‑Experts model locally on an 8 GiB NVIDIA GeForce GTX 1080 using an enhanced llama.cpp fork (AtomicBot-ai/atomic-llama-cpp-turboquant). The guide details system setup (driver pinning, CUDA nvcc, gcc-14 workaround, glibc patch), building the fork with CUDA, downloading the main GGUF and MTP assistant head, and tuning offload parameters. Key optimisations include keeping most MoE expert weights in host RAM (streamed over PCIe), using RotorQuant/TurboQuant KV cache to enable 128k context, and forcing the assistant embedding table onto the GPU with --override-tensor-draft to enable effective MTP speculative decoding. The author reports ~24.5 tokens/sec at 128k context and describes the memory/PCIe tradeoffs and the final recommended command-line configuration.

Read assessment

Track Real-Time Market Signals & Shifts

Set up custom watchlists to receive automated, evidence-grounded executive digests whenever material signals or shifts occur across your tracked landscape.