DEV Community

xbill for AWS Community Builders

Posted on

Pure JAX on G5g: Serving Gemma 4 on Graviton and a T4G

This article provides a step by step deployment guide for serving Google's Gemma 4 on an AWS EC2 G5g instance using pure JAX.

The code is here:

github.com/xbill9/gemma4-dev

What is this project trying to Do?

This project aims to serve a modern open model on the cheapest whole CUDA GPU AWS will rent you, and to measure honestly what that costs.

Aren't You Using The Wrong GPU?

Probably! The T4G is a Turing chip from 2018. It has no bfloat16 and no fp8.

But it is cheap, it is available when nothing else is, and it is attached to a Graviton2 host — which makes G5g the rare hardware axis that almost nothing in the ML ecosystem targets: aarch64 and CUDA together.

So let's give pure JAX a shot on G5g!

AWS EC2 G5g

G5g instances pair an AWS Graviton2 (64-bit Arm) processor with NVIDIA T4G Tensor Core GPUs. At g5g.xlarge they are the cheapest EC2 instance carrying a whole NVIDIA GPU, and the only Arm-based GPU family AWS offers.

Two GPU instances are cheaper per hour and neither can serve this model (us-east-1, Linux, on-demand, checked against the Pricing API on 2026-08-28):

  • g6f.large at $0.2020 is genuinely NVIDIA and genuinely CUDA — but it is one eighth of a GPU with 3 GB, and the weights alone are 6.155 GB. The first g6f that fits is g6f.4xlarge at $0.9500, which is 1.7x this rig's g5g.2xlarge.
  • g4ad.xlarge at $0.3785 carries an AMD Radeon Pro V520 — no CUDA at any price.

Among whole NVIDIA GPUs, G5g is the floor: g5g.xlarge at $0.4200, and the next one up is g4dn.xlarge at $0.5260.

More information is available here:

https://aws.amazon.com/ec2/instance-types/g5g/

The default in this rig is g5g.2xlarge — 1 GPU, 8 vCPU, 16 GiB RAM.

Note- the T4G reports 15,360 MiB of device memory, not the nominal 16 GB. Budget against the measured number.

Gemma 4

Gemma is Google's family of open models built from the same research as Gemini. This rig serves google/gemma-4-E2B-it, the instruction-tuned reference release.

JAX

JAX is Google's array computing library — NumPy semantics, composable transformations, and compilation to XLA. On NVIDIA hardware, pip supplies the CUDA libraries, so there is nothing to build.

More information is available here:

https://github.com/jax-ml/jax

"Pure JAX" here is literal. The engine is this repo's own Gemma 4 port driven by a JAX generation loop behind an OpenAI-compatible FastAPI server, running under systemd.

Prerequisites

You need four things before starting:

  • An AWS account with G instance vCPU quota in us-east-1
  • A subnet id, a security group id, and an IAM instance profile — the rig requires all three explicitly and will not create them for you
  • A Hugging Face token with access to the Gemma weights
  • Python 3.13 and pip

The instance profile needs AmazonSSMManagedInstanceCore plus read access to your Secrets Manager secret and your S3 cache bucket. There is no inbound SSH rule and no private key — all remote administration goes over SSM Run Command.

Install the Rig

Clone the monorepo and install the control plane:

git clone https://github.com/xbill9/gemma4-dev
cd gemma4-dev/gpu-jax-g5g-2b
pip install -r requirements.txt
Enter fullscreen mode Exit fullscreen mode

That installs boto3 and FastMCP only. Nothing here needs a GPU — the GPU is on the other end.

Run the Tests

python3 -m unittest discover -s tests -v
Enter fullscreen mode Exit fullscreen mode

105 tests, fully offline. Every cloud, subprocess, and network boundary is mocked. If these do not pass, do not launch an instance.

Register the MCP Server

The whole rig is driven by an MCP server exposing a devops agent:

./project-setup.sh
Enter fullscreen mode Exit fullscreen mode

This installs the bundled skill and registers .mcp.json:

{
  "mcpServers": {
    "gpu-jax-g5g-2b": {
      "command": "python3",
      "args": [".claude/skills/gpu-jax-g5g-2b-management/mcp/server.py"],
      "env": {
        "AWS_REGION": "us-east-1",
        "MODEL_NAME": "google/gemma-4-E2B-it",
        "INSTANCE_TYPE": "g5g.2xlarge",
        "MCP_SERVER_NAME": "gpu-jax-g5g-2b"
      }
    }
  }
}
Enter fullscreen mode Exit fullscreen mode

Every tool is now available as mcp__gpu-jax-g5g-2b__<tool>.

Save the Hugging Face Token

save_hf_token(token="hf_...")
Enter fullscreen mode Exit fullscreen mode

This writes to AWS Secrets Manager under vllm/hf-token. The instance fetches it at boot into a root-only EnvironmentFile.

Note- the token never goes in user data. Instance metadata is readable by anything running on the box.

Check the Quotas

check_g5g_quotas()
Enter fullscreen mode Exit fullscreen mode

Reports your On-Demand and Spot G instance vCPU limits for the region. g5g.2xlarge is 8 vCPU. Check this before launching, not after the launch fails.

Launch the Instance

create_g5g_instance(
    subnet_id="subnet-...",
    security_group_id="sg-...",
    iam_instance_profile="...",
    spot=True
)
Enter fullscreen mode Exit fullscreen mode

The AMI is resolved at launch time from SSM Parameter Store:

/aws/service/deeplearning/ami/arm64/base-oss-nvidia-driver-gpu-ubuntu-26.04/latest/ami-id
Enter fullscreen mode Exit fullscreen mode

Never hardcode an AMI id here. AWS also ships ARM64 DLAMIs built for Graviton CPU inference. They boot perfectly and simply have no GPU. The /latest/ parameter also moves — this rig has seen ami-0bff4343bfd56a20e become ami-025a6e5b3b786cf61 overnight as Ubuntu 26.04 became 26.04.1.

Watch the Install

get_install_progress(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

Cloud-init installs jax[cuda13] on Python 3.14. The stages are timed:

[stage] jax-wheels      43s   (total  84s)
[stage] serving-deps    14s   (total  98s)
[stage] gpu-verify      13s   (total 111s)
[stage] cache-restore    6s   (total 117s)
[stage] unit-rewrite     0s   (total 117s)
Enter fullscreen mode Exit fullscreen mode

117 seconds, and there is no compile step anywhere in it. That is the entire reason this rig exists — more on that below.

The cache-restore stage pulled 805 files / 12 MB in 6 seconds from S3, on a fresh instance, compiled by a box that had already been terminated. XLA's cold-compile penalty becomes a rounding error on Spot.

Verify the GPU

verify_gpu_arch(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

This measures whether JAX's CUDA kernels actually cover this GPU, rather than trusting that they do. You want to see SM 7.5 claimed and a real device, not a silent CPU fallback.

Deploy the Server

make skill
deploy_jax_server(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

Always make skill first. The deploy ships the skill snapshot, not your working tree. The deploy output prints the payload root and the build id so a stale deploy is visible in one line.

Verify The Installation

The first line the process emits is the device-policy banner:

INFO ports.gemma4.jax_e_model: jax_e_model device policy: platform=gpu
compute_capability=7.5 compute_dtype=float16 pallas_interpret=False
Enter fullscreen mode Exit fullscreen mode

Both halves matter. float16 is the device choosing Turing's only real 16-bit datapath — it is read from the live compute capability, not from a config file. And pallas_interpret=False is the difference between serving and silently running a simulator.

Then the whole resolved configuration lands on one greppable line:

READY build_id=6852f5680f43 ... compute_dtype=float16 kv_cache_dtype=float16
kv_cache_requested=auto pre_ampere=True quant_mode=fp16 window_kv=True
Enter fullscreen mode Exit fullscreen mode

Load is staged, so a hang is attributable:

download        87.7s
read_shards     73.5s   (1 shard, 600 tensors, 0.95 GB of non-text towers skipped)
convert_params   3.4s
device_put       0.0s
                164.7s total, 9.26 GB
Enter fullscreen mode Exit fullscreen mode

Then confirm the served build matches what you shipped:

verify_model_health(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

Query the Model

query_model(instance_id="i-...", prompt="Explain Graviton in one sentence.")
Enter fullscreen mode Exit fullscreen mode

The endpoint is OpenAI-compatible on :8000, so anything that speaks that API works:

get_endpoint(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

Note- warm up at the shape you measure. max_new_tokens is a static_argnames entry, so (bucket, max_tokens) is the compiled shape. The same request measured 18.77 s cold against 4.35 s warm.

Check the Metrics

get_metrics(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode
tpu_jax_decode_tokens_per_second   13.0
tpu_jax_prefill_milliseconds      692.6
tpu_jax_hbm_used_bytes            6296892160
tpu_jax_weight_bytes              6155450950
tpu_jax_degenerate_responses_total   0
Enter fullscreen mode Exit fullscreen mode

Quote the gauge, not end-to-end. Decode is flat at 12.9 / 13.0 / 12.9 tok/s across 41 → 2,057 input tokens. End-to-end throughput does fall (12.43 → 8.22) — but that is prefill being linear in the padded bucket, not decode degrading. Two different claims.

Why Not vLLM?

Because I tried, on identical silicon, and it works — at 43 tok/s — but only after this:

  1. A ~67-minute from-source build
  2. cuda-toolkit from NVIDIA's sbsa repo, because the DLAMI ships a driver but no nvcc
  3. A Rust toolchain, for vLLM's frontend
  4. An unlanded patch to Triton's attention kernel, reapplied after every upgrade

That last one is the interesting failure. Gemma 4 has heterogeneous attention head dimensions — sliding layers at 256, global layers at 512. Only two vLLM backends support that, and with FA4 unavailable it force-selects Triton:

Gemma4 model has heterogeneous head dimensions
{'sliding_attention': 256, 'full_attention': 512}. FA4 not available,
forcing TRITON_ATTN backend.
Enter fullscreen mode Exit fullscreen mode

Whose 512-wide tile then asks Turing for memory Turing does not have:

triton.runtime.errors.OutOfResources: out of resource: shared memory,
Required: 98304, Hardware limit: 65536
Enter fullscreen mode Exit fullscreen mode

JAX sidesteps all four. pip supplies CUDA, so no build, no toolkit, no Rust. The plugin's precompiled cubins already cover sm_75. And attention is ordinary XLA rather than a hand-tiled Triton kernel, so there is no per-block shared-memory ceiling and no patch to carry.

The honest trade: 13.10 tok/s against 43, for a 117-second install with nothing to reapply. For a measurement rig I re-provision constantly on Spot, that was right. For a production endpoint it probably is not.

The Number I Did Not Expect

I profiled decode with xprof. The kernel table is the whole story:

conversion   54.0%   <-- dtype conversion
fp32 gemv    32.9%
fusion       12.2%
TensorCore    0.0%
Enter fullscreen mode Exit fullscreen mode

Zero. 1,466 ms of kernels across 108 distinct kernels on a Tensor Core GPU, and not one Tensor Core fired. Over half of decode went to converting numbers between formats before the math could start.

The obvious hypothesis was bf16 weights being converted on a chip with no bf16. So this weekend I converted the checkpoint to float16 host-side and re-ran. Parameter dtypes now read {'float16': 541, 'uint8': 1, 'int8': 1} — and conversion is still 54.0%.

The obvious explanation is wrong and I do not yet know the real one. That is the next thing to profile. I would rather publish the open question than a tidy story.

What I can stand behind is that the measurement is real: 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.

AWS Services Used

Service Role
EC2 (g5g.2xlarge, Spot) Graviton2 + NVIDIA T4G
Systems Manager — Parameter Store Resolves the arm64 GPU DLAMI id at launch
Systems Manager — Run Command Ships the payload, runs every diagnostic. No SSH
S3 XLA compilation cache, shared across instances
Secrets Manager Hugging Face token
IAM Instance profile scoping all of the above
Service Quotas Pre-flights G vCPU limits before a doomed launch
EBS gp3 100 GB at 500 MiB/s, 6,000 IOPS, for a 9.5 GB checkpoint

Everything goes through boto3. The rig never shells out to the AWS CLI.

Clean Up

terminate_g5g_instance(instance_id="i-...")
Enter fullscreen mode Exit fullscreen mode

Termination is cheap on this rig — there is no built image to lose with the root volume, only a pip install and a model cache, and the compilation cache is already in S3.

What I Learned This Summer

Check reachability on paper before spending a provisioning cycle. A nine-step analysis order — does the compute dtype match the chip, is there a fused kernel for that format on that chip — runs in an afternoon. A launch costs a day. It has killed two bad plans before either touched hardware.

A wrong dtype does not error. It emulates. bfloat16 on Turing does not fail loudly; it routes through fp32 and quietly eats your decode.

Refuse early, with the arithmetic attached. The fused W4A16 kernel wants 550 KiB–1.1 MiB per block and Turing gives you 64 KiB. The rig computes that at startup and refuses with the numbers in the message, rather than dying as a cryptic OutOfResources at the first token.

The scariest bugs return status: "success". A padding-eviction bug in the KV ring cache produced a token loop, not a crash. Nothing in the logs was red. It took a week.

Summary

The AWS G5g instance provides a genuinely cheap environment for serving open models, and pure JAX reaches a served token on it without a single line of compiled code. The throughput is not competitive with a patched vLLM - but the deployment is 117 seconds, reproducible to 1.1 ms, and has nothing to reapply.

Top comments (0)