Porting Gemma-4 to AWS Inferentia2 Hardware
- •Gemma-4 models ported to AWS Inferentia2 by bypassing vendor-standard inference stacks
- •Porting requires manual sharding logic for cross-layer KV-sharing and mixed attention head types
- •Compiler constraints necessitate host-side embedding storage and device-level logit optimization
Running Google's Gemma-4 model family—specifically the 2B, 4B, and 12B versions—on AWS Inferentia2 (inf2) hardware requires bypassing the standard vendor inference path to address architectural incompatibilities. The AWS vendor stack, including optimum-neuron and the Neuron vLLM backend, does not natively support Gemma-4's specific features, such as per-layer embeddings, cross-layer key/value (KV) sharing, and mixed attention types. Developers must instead use direct tracing via torch_neuronx and neuronx_distributed to bypass these vendor abstractions.
Porting the models involves resolving three distinct types of 'mixed heads' issues. First, E2B and E4B models utilize cross-layer KV-sharing, where shared layers do not compute their own projections. Standard graph builders fail to represent this, but by tracing the Hugging Face forward pass directly, KV-sharing becomes a live data dependency. Second, E4B models face a Grouped-Query Attention (GQA) mismatch under tensor parallelism (TP=2) because the KV head count (2) does not divide evenly by the rank requirements when attempting naïve sharding. Correcting this requires keeping num_key_value_groups unchanged to maintain the correct mapping between query and KV heads.
Third, the 12B model interleaves sliding-window and global attention layers with differing KV-head counts, which creates an indivisible sharding problem. The solution involves leaving global layers' KV replicated while shrinking group counts to match sharded query-head counts. Additionally, the Neuron compiler (neuronx-cc) imposes a 196,608-byte limit per partition for the on-chip state buffer (SBUF). To avoid overflow, the per-layer embedding table must remain on the host, and logit soft-capping is disabled on the device since greedy decoding is unaffected by the change. These optimizations allow E2B to achieve approximately 44 tokens per second (tok/s) on one core, while E4B and 12B reach 33–39 tok/s and 15 tok/s, respectively, on TP=2 configurations.