Back to blog

Gemma 4 Optimization with Agentic Kernel Engineering

Sep 29, 2026 · 6 min read

OpenRelay is a hardware-agnostic inference platform. We run models on NVIDIA and AMD alike, and we are a little obsessive about one thing: getting every TFLOP and every gigabyte of HBM a chip has into the workloads our customers run, at the latency they hold us to.

A spec sheet tells you what a GPU could do. What it actually does depends on a pile of kernels, lookup tables and flags, most of them tuned for some other model. Closing that gap is daily work for us, and these days agents do most of it. Here is one evening of that work: Gemma 4 31B on AMD MI355X.

1.5x
traffic per GPU at the same latency target
2.0 to 3.0 req/s
22 ms
per output token at that load
from 40 to 43 ms
41.9K
prefill tokens/s per GPU
from 27.8K
98.8%
GSM8K, before and after
no accuracy traded
1 s2 s4 s8 s16 s32 s1.52.02.53.03.5Requests per second, per GPUTarget: p95 under 4 s1.5x the traffic, same targetFirst config, 1.75 req/s per GPU: p95 TTFT 2.24 s and 2.54 s (two seeds), median TPOT 27.9 and 27.2 msFirst config, 2 req/s per GPU: p95 TTFT 2.9 s and 2.85 s (two seeds), median TPOT 43.1 and 39.9 msFirst config, 2.25 req/s per GPU: p95 TTFT 5.75 s and 5.62 s (two seeds), median TPOT 65.3 and 66.1 msFirst config, 2.5 req/s per GPU: p95 TTFT 25.9 s and 27.4 s (two seeds), median TPOT 69.4 and 69.7 msBeforeFirst config (2.0 req/s under target)After one evening, 2.75 req/s per GPU: p95 TTFT 1.59 s and 1.42 s (two seeds), median TPOT 17.3 and 15.1 msAfter one evening, 3 req/s per GPU: p95 TTFT 1.67 s and 1.8 s (two seeds), median TPOT 21.6 and 22.6 msAfter one evening, 3.25 req/s per GPU: p95 TTFT 3.37 s and 4.49 s (two seeds), median TPOT 36.1 and 42.1 msAfter one evening, 3.5 req/s per GPU: p95 TTFT 10.7 s and 12.8 s (two seeds), median TPOT 45.6 and 46.7 msAfterAfter one evening (3.0 req/s under target)
One MI355X, Gemma 4 31B. The latency target is p95 time to first token under 4 s. Open-loop load, 300 s per rate, the worse of two seeds, log scale.
Data
ConfigReq/s per GPUp95 TTFT (seed 7 / 8)Median TPOT (seed 7 / 8)
First config1.752.24 / 2.54 s27.9 / 27.2 ms
First config22.9 / 2.85 s43.1 / 39.9 ms
First config2.255.75 / 5.62 s65.3 / 66.1 ms
First config2.525.9 / 27.4 s69.4 / 69.7 ms
After one evening2.751.59 / 1.42 s17.3 / 15.1 ms
After one evening31.67 / 1.8 s21.6 / 22.6 ms
After one evening3.253.37 / 4.49 s36.1 / 42.1 ms
After one evening3.510.7 / 12.8 s45.6 / 46.7 ms

A new chip, the same model

When we moved Gemma 4 31B from NVIDIA B200 to AMD MI355X, getting it running took a day. Blackwell's NVFP4 weights don't run on MI355X, so we switched to AMD's Quark MXFP4 checkpoint and kept Google's speculative-decoding drafter. Memory was the easy part: weights and drafter come to about 20 GB, which leaves most of the chip's 288 GB of HBM for KV cache, and long prompts eat KV cache.

Compute was another story. The first working config did 27.8K prefill tokens per second per GPU, which sounds fine until you do the arithmetic: the matrix math was running at under a fifth of the chip's MXFP4 peak. Held to our latency target, p95 time to first token under four seconds, one GPU topped out at 2.0 requests per second.

Agentic kernel engineering

This is the kind of gap we hand to agents. The loop they run is deliberately boring.

Research agentReads kernel source and open PRs,ranks the likely leversJournalEvery command, number andmistake, written downProfileWhere does theGPU time go?HypothesisBiggest slicefirstExperimentSmallest changethat could failGateGSM8K ≥ 96.5%or it doesn't countKeep or revertSame GPU,same harnessnext slice of the profile, until the hardware says stop
The loop is deliberately boring. The gate is the step we never skip: if a faster kernel changes the model's answers, we treat it as a bug.

A research agent reads the kernel libraries' source and open pull requests and ranks what is likely to matter. A tuning agent works on real GPUs, one change at a time, and nothing it measures counts unless the model still scores at least 96.5% on GSM8K. It writes everything down, mistakes included. It caught one this time: a round of benchmarks that had quietly measured the wrong engine. It threw those numbers out and fixed its launcher so that can't happen again.

For Gemma 4 we gave it four MI355X GPUs and an evening.

25K30K35K40K45K21:0021:3022:0022:3023:0023:3000:0000:30first profilecompile storm fixedlatency runs donefirst configtuned GEMM rowsfused GELU + quantgraphs + split-KVnative RMSNorm8K prefill chunkattention table entries: +33%41.9K
Best prefill throughput per GPU as the evening went on, 16 concurrent 10K-token prompts. Times from the agent's journal. The biggest step was two new entries in a lookup table.
Data
TimeChangePrefill tokens/s
21:10first config27.8K
21:15attention table entries37.1K
21:33tuned GEMM rows38.7K
21:48fused GELU + quant39.4K
22:18graphs + split-KV39.8K
23:05native RMSNorm41.6K
23:358K prefill chunk41.9K

What it found

Seven minutes in, the first profile came back with attention at 43% of prefill time. AMD's kernel library, AITER, picks attention tile sizes from a lookup table per GPU, and the MI355X table had no prefill entry for Gemma 4's large attention heads. So prefill ran on a tile sized for decoding, which re-read the same keys and values over and over.

K/V CACHE FOR ONE KV HEADBefore: 16-row tile2 tokens per program,8 programs8 passesAfter: 128-row tile16 tokens, 1 program1 passsame keys and values, same answer
Illustration, not a trace. For the same 16 tokens of one global-layer KV head, the decode-sized tile read every K/V tile eight times. The prefill-sized tile reads it once.

Two new entries in that table made attention 2.6x faster and prefill a third faster, with the model's answers unchanged.

AttentionMXFP4 matmulsEverything elseBefore4824501891,121 msAfter185453104742 ms
GPU time in milliseconds to prefill the same three 10K-token prompts. Attention fell 2.6x. The matrix multiplies barely moved and are now 61% of the work, which is where you want a GPU spending its time.
Data
AttentionMXFP4 matmulsEverything elseTotal
Before482 ms450 ms189 ms1,121 ms
After185 ms453 ms104 ms742 ms

The second find was a GPU waiting on its CPU. With speculative decoding, each request needs three tokens per decode step, and the engine had only captured CUDA graphs up to 64 tokens. Any batch over 21 requests launched its kernels one at a time from the CPU. One flag fixed it.

Kernels runningGPU idle, waiting on the next kernel launchGraphs to 64 tokens21.3 ms36.4 ms idle57.7 msGraphs to 256 tokens21.3 ms2.7x shorter step, GPU busy the whole time
One decode step with about 30 requests in flight. Before, the GPU spent 63% of the step waiting for the CPU to launch the next kernel (summed here; in the trace the gaps sit between launches).
Data
Graph capture sizesKernel timeStep wall timeGPU busy
up to 64 tokens21.3 ms57.7 ms37%
up to 256 tokens21.3 ms21.3 ms100%

The rest were smaller: tuned matrix-multiply configs for Gemma's exact shapes, a fused activation-and-quantize kernel, native RMSNorm so the compiler can fuse around it, and a kernel that was recompiling itself for every new batch size in the middle of serving, 190 milliseconds at a time. By the end of the evening, one MI355X served 1.5x the traffic under the same latency target, tokens streamed at 22 ms instead of 40 to 43, and GSM8K read 98.8% before and after.

How much of the chip we use now

Matrix math is now 61% of prefill time, which is where you want a GPU spending it, and those kernels run at 55 to 61% of the MXFP4 peak at the clocks the chip holds under load. Better kernels might buy another 1.2 to 1.3x. Past that, physics wins.

MeasuredEstimateMXFP4 peak: matmuls alone top out near 170KFirst working config27.8KAfter one evening41.9KBetter kernels (our estimate)50K to 55KIf attention and small ops were free68KIf matmuls hit peak at sustained clocks125K10x the first config278K
How much of the chip we use, in prefill tokens/s per GPU. Solid bars are measured; hatched bars are estimates from the profile and the chip's MXFP4 peak. Past the wall, no kernel helps.
Data
Prefill tokens/sKind
First working config27.8Kmeasured
After one evening41.9Kmeasured
Better kernels (our estimate)50K to 55Kestimate
If attention and small ops were free68Kestimate
If matmuls hit peak at sustained clocks125Kestimate
10x the first config278Kestimate

We are open-sourcing all of it

Everything the agent changed is in OpenRelayInc/inference-recipes: the Dockerfile, the tuned tables, the patches, and the benchmark harness behind every number in this post. The fixes that belong upstream are pull requests to AMD's AITER and to vLLM, each re-measured on their current main branch.

Where each fix lands in the open serving stack, from the engine down to the kernels. Gains are re-measured on each project's current main branch on an MI355X.
Data
Pull requestLayerChangeOn main
vllm-project/vllm#59153vLLMfused activation + MXFP4 quant+1.8% prefill
ROCm/aiter#5926Attention kernelsprefill tile for 512-dim heads2.7x
ROCm/aiter#5929Attention kernelssplit-KV for speculative decodingup to 3.8x
ROCm/aiter#5928GEMM tuning tablesGEMM rows tuned for Gemma 4up to 1.8x
ROCm/aiter#5927Activation and quant kernelsno recompile per batch size12 → 4 compiles

One of our fixes didn't make the cut. AITER's main branch has a new MI355X attention kernel that already runs Gemma 4's sliding-window layers faster than our tuned tile did (0.55 ms against our 0.61), so we left that one out and cheered instead.

For the kernel engineers

How we measured

One MI355X, the same harness before and after. The baseline is our first working MI355X config on stock vllm/vllm-openai-rocm:v0.30.0, not B200. We never ran this harness on B200, so nothing here compares vendors. Load tests are open-loop, 300 seconds per rate, two seeds, with requests averaging 9.5K tokens in and 380 out. Prefill is measured with 16 concurrent 10K-token prompts.

The tile

Gemma 4 has 50 sliding-window layers with 256-dim heads and 10 global layers with 512-dim heads and 8 query heads per KV head. On the global layers a 16-row tile is 2 tokens per program. The entry, keyed on long queries so decode keeps its tile:

"D_GEQ_512.Q_GEQ_256.DT_fp8_fp8": {"BLOCK_M": 128, "TILE_SIZE_MIN": 64, "TILE_SIZE_MAX": 64,
                                   "num_warps": 4, "num_stages": 1, "waves_per_eu": 1}

Graphs and the compile storm

Two draft tokens make each decode request 3 tokens, so capture sizes that stop at 64 leave batches over 21 requests ungraphed. Capturing to 256 costs 7 seconds of startup and 3.1 GiB. Separately, AITER declared a padded row count (scaleM_pad) as a Triton constexpr, so every new batch size compiled a new kernel mid-serving. As a runtime argument, 14 batch sizes need 4 compiles instead of 12, with bit-identical output.

First config+ attention tile entries+ tuned MXFP4 GEMM rows+ CUDA graphs to 256 tokens+ fused GELU, up, MXFP4 quant+ split-KV decode attention+ native RMSNormPrefill tokens/s16 concurrent 10K-token prompts27.9K37.3K +34%38.7K38.8K39.4K39.5K41.7KRequests/s32 users, 9.5K in, 380 out2.082.322.422.88 +19%2.942.983.05
Each row adds one change to the row above, on the same GPU. The tile entry moved prefill; the graph fix moved end-to-end throughput. The shipped config adds an 8,192-token prefill chunk to the last row.
Data
StagePrefill tokens/sRequests/s
First config27.9K2.08
+ attention tile entries37.3K2.32
+ tuned MXFP4 GEMM rows38.7K2.42
+ CUDA graphs to 256 tokens38.8K2.88
+ fused GELU, up, MXFP4 quant39.4K2.94
+ split-KV decode attention39.5K2.98
+ native RMSNorm41.7K3.05

What didn't help: Triton and hipBLASLt MXFP4 GEMMs (2.5x slower than AITER's assembly kernels), vLLM's QK-norm and RoPE fusion passes (they don't match Gemma 4's layout), CUDA graphs up to 16K tokens (no gain for 14 GiB), 128 concurrent sequences (twice the time per token for 1 to 3% more throughput), and a third draft token (same capacity).

Try it

Three of the changes are flags on the stock image. On vLLM after v0.30.0, also set VLLM_ROCM_USE_AITER_FP4_ASM_GEMM=1: without it, MXFP4 GEMMs fall back to Triton and Gemma 4 prefill drops from 27.7K to 18.9K tokens per second on a recent nightly.

--max-num-batched-tokens 8192
--compilation-config '{"cudagraph_capture_sizes":[1,2,4,8,16,24,32,48,64,80,96,112,128,144,160,176,192,224,256]}'
--kernel-config '{"ir_op_priority":{"rms_norm":["native"]}}'

Run Gemma 4 31B on OpenRelay

One model ID, on whichever GPU serves it best.

If you would like to build the agents that do this every day, we are hiring.