unsloth/scripts/benchmarks/gemma4_flex_inference.py
Daniel Han 8fb0c2e2a7 benchmarks: add gemma4_flex_inference for unsloth/gemma-4-E2B-it
Extends the flex_attention + paged KV + CUDA graphs engine from Qwen3 and
Llama-3.2 to Gemma-4-E2B-it via a new standalone file that imports the
shared helpers (PagedKVCache, PageTable, Sequence, LoRA double-copy,
drift verification, flex_attention_compiled, _apply_rotary) from
qwen3_flex_inference.py. The Qwen3 / Llama path is not modified.

Gemma-4 diverges from Qwen3 / Llama in ways that cannot be folded into a
single hasattr guard:

- KV-sharing layers. E2B has 35 layers; the upper 20 lack k_proj / v_proj
  / k_norm / v_norm and consume the full prefix K/V produced by a store
  layer further up the stack. We allocate a sidecar dict of
  [max_batch, n_kv, max_seq, head_dim] buffers at fixed device addresses,
  populated by store layers during prefill and read by shared layers
  through eager SDPA (their layout does not match the paged cache's
  block-mask shape).
- Dual attention regimes. full_attention (head_dim=512, rope_theta=1e6)
  and sliding_attention (head_dim=256, sliding_window=512) coexist. We
  precompute both (cos, sin) pairs via
  Gemma4TextRotaryEmbedding(x, position_ids, layer_type) and dispatch on
  self.layer_type inside the patched attention forward.
- Per-layer input embeddings. The walker threads the
  [B, S, num_layers, hidden_size_per_layer_input] table from
  get_per_layer_inputs + project_per_layer_inputs through each layer's
  per_layer_input_gate / act_fn / mul / per_layer_projection /
  post_per_layer_input_norm path.
- Four norms per block with double residuals. Attn (input_layernorm,
  post_attention_layernorm) and MLP (pre_feedforward_layernorm,
  post_feedforward_layernorm), plus the per-layer-input residual and
  layer_scalar multiply.
- Final logit softcap. tanh(logits / 30.0) * 30.0 on the lm_head output.

Transformers>=5.5.0 is required for the gemma4 module. A _require_gemma4
guard at main() exits with a clear install hint when the module is
missing, so the workspace's Qwen3 / Llama path stays on the existing
transformers install.

The text-only path loads Gemma4ForConditionalGeneration, drops the
vision and audio towers, and moves the language_model into a
Gemma4ForCausalLM shell so LoRA, state-dict hashing, and the double-copy
refresh treat it like any other HF decoder model.

CLI mirrors qwen3_flex_inference.py: --model_name (default
unsloth/gemma-4-E2B-it), --lora_adapter, --load_in_4bit,
--capture_cudagraph, --verify_no_drift, --chat_template {auto,grpo,
native}, --fa4_prefill / --no-fa4_prefill, plus decode / prefill
kernel_options for Triton block tuning.

Smoke-tested on B200 (sm_100) with bf16, batch 2, max_new_tokens 16,
no-fa4_prefill + BLOCK_M=32 / BLOCK_N=32 (Gemma-4 head_dim=256 exceeds
FA4's 128-limit on sm_100): 52 tok/s cold, coherent completions.

scripts/benchmarks/README.md gains a paragraph covering the
transformers>=5.5 dependency, the head_dim=256 constraint, and the
recommended kernel_options for B200.
2026-04-21 08:45:05 +00:00

1071 lines
41 KiB
Python

"""Gemma-4-E2B-it inference with flex_attention + paged KV cache + CUDA graphs.
Extends the Qwen3/Llama-3.2 engine in `qwen3_flex_inference.py` to a third
architecture, `unsloth/gemma-4-E2B-it`. Gemma-4 is not a drop-in addition:
its text backbone diverges from Qwen3/Llama in ways that cannot be folded
into a single `hasattr(self, "q_norm")` branch. The divergences, and how
this file handles them:
1. KV-sharing layers. E2B has 35 layers; the upper 20 (layers 15-34) lack
`k_proj`, `v_proj`, `k_norm`, `v_norm` entirely and consume the full
prefix K/V produced by a "store" layer further up the stack. Paging
the KV for shared layers would cost more than it saves, so the shared
layers read from a pre-sized sidecar dict
(`FlexGemma4Inference.shared_kv_buffer`) whose tensors live at fixed
device addresses for CUDA graph safety.
2. Dual attention types. Each layer is either `full_attention`
(`head_dim=512`, `rope_theta=1e6`, `partial_rotary_factor=0.25`) or
`sliding_attention` (`head_dim=256`, `rope_theta=10000`,
`sliding_window=512`). We precompute both (cos, sin) pairs once per
forward and dispatch on `self.layer_type`.
3. Per-layer input embeddings. `embed_tokens_per_layer` produces a
`[B, S, num_layers, 256]` auxiliary table that enters every layer
through a `per_layer_input_gate -> act -> mul -> per_layer_projection
-> post_per_layer_input_norm -> +residual` path after the MLP residual.
4. Four norms per layer. `input_layernorm` / `post_attention_layernorm`
wrap the attention block (double residual); `pre_feedforward_layernorm`
/ `post_feedforward_layernorm` wrap the MLP (double residual). A scalar
`layer_scalar` multiplies hidden_states at layer end.
5. Final logit softcap. `logits = tanh(logits / 30.0) * 30.0` applied on
the lm_head output.
The engine is text-only: `Gemma4ForCausalLM(text_config)` skips the
multimodal `Gemma4ForConditionalGeneration` wrapper and its vision + audio
towers entirely. Shared helpers (`PagedKVCache`, `PageTable`, `Sequence`,
`refresh_lora_merge_from_pristine`, `run_drift_verification`,
`flex_attention_compiled`, `_apply_rotary`, FA4 capability guard) are
imported from `qwen3_flex_inference.py` unchanged.
Run:
CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/gemma4_flex_inference.py \
--n_prompts 64 --max_new_tokens 512 --capture_cudagraph \
--stats_path logs/flex_gemma4_bf16.json
Requires `transformers>=5.5.0` (the `gemma4` module). The main workspace
env stays on 4.57.6; this file short-circuits with a clear install hint
if the module is missing. Use `isolated_run.py` with
`--extra_packages "transformers>=5.5.0 peft datasets"` to run on that env.
"""
from __future__ import annotations
import argparse
import copy
import json
import os
import sys
import time
import types
from collections import deque
from pathlib import Path
from typing import Optional
import torch
import torch.nn.functional as F
from torch.nn.attention.flex_attention import BlockMask
HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))
# Shared helpers from qwen3_flex_inference.py. Import-time cost is paid once;
# we do not re-define any of these locally.
from qwen3_flex_inference import ( # noqa: E402
DECODE_KERNEL_OPTIONS_DEFAULT,
PREFILL_KERNEL_OPTIONS_DEFAULT,
Sequence,
_apply_rotary,
_hash_state_dict,
_lora_needs_peft_fallback,
flex_attention_compiled,
refresh_lora_merge_from_pristine,
run_drift_verification,
)
from flex_paged_attention import PagedKVCache, PageTable # noqa: E402
# --- transformers version guard --------------------------------------------
#
# Gemma-4 lands in `transformers>=5.5.0`. The main workspace env is on
# 4.57.6, which keeps the Qwen3 + Llama paths in qwen3_flex_inference.py
# working unchanged. We defer the import to call time so `--help` still
# works on 4.57.6.
def _require_gemma4():
try:
from transformers.models.gemma4.modeling_gemma4 import Gemma4ForCausalLM
from transformers.models.gemma4.configuration_gemma4 import (
Gemma4Config,
Gemma4TextConfig,
)
except ImportError:
import transformers
raise SystemExit(
f"Gemma-4 requires transformers>=5.5.0 (`gemma4` module). "
f"Current: transformers=={transformers.__version__}. "
f"Install: uv pip install 'transformers>=5.5.0'"
)
return Gemma4ForCausalLM, Gemma4Config, Gemma4TextConfig
# --- attention forward factory --------------------------------------------
def _apply_rotary_q(q, cos, sin):
"""Rotary on Q alone; used when K comes from the shared sidecar and
already carries its rotary from the store layer."""
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
def rotate_half(x):
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim = -1)
return (q * cos) + (rotate_half(q) * sin)
def make_flex_gemma4_attention_forward(
page_table: PageTable, shared_kv_buffer: dict
):
"""Return a new `forward` method for `Gemma4TextAttention` that routes
through flex_attention against either the paged cache (non-shared
layers) or the full-length shared KV sidecar.
Three layer kinds:
- shared (`self.is_kv_shared_layer == True`): no `k_proj`/`v_proj`/
`k_norm`/`v_norm`. Read K/V from
`shared_kv_buffer[self.kv_shared_layer_index]`.
- store (`self.store_full_length_kv == True`): standard q/k/v
projection. After rotary, write the full-sequence K/V
into `shared_kv_buffer[self.layer_idx]` on prefill.
Also updates the paged cache for its own attention.
- plain (neither flag set): standard q/k/v + paged cache.
Shared layers attend over the prefix K/V only (populated at prefill).
Decode-time tokens generated by store layers are NOT written into the
sidecar -- this is the documented simplification from the plan; it
trades exactness during decode for a fixed-address sidecar that is
safe under CUDA graph capture.
`position_embeddings` is a dict keyed by `layer_type`; we pick the
right (cos, sin) pair before rotary.
Expects `self._paged_cache` to be set on non-shared layers (None on
shared layers). Shared layers still keep `self.q_proj`, `self.q_norm`.
`self.v_proj` may be None when the config sets `attention_k_eq_v`
(global head dim with shared K=V); that branch reuses k_states.
"""
def forward(
self,
hidden_states: torch.Tensor,
position_embeddings: dict,
attention_mask = None,
past_key_values = None,
cache_position = None,
flex_block_mask: Optional[BlockMask] = None,
flex_input_pos: Optional[torch.Tensor] = None,
flex_batch_idx: Optional[torch.Tensor] = None,
flex_kernel_options: Optional[dict] = None,
**kwargs,
):
input_shape = hidden_states.shape[:-1]
hidden_shape = (*input_shape, -1, self.head_dim)
cos, sin = position_embeddings[self.layer_type]
q = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
if getattr(self, "is_kv_shared_layer", False):
# Shared layer: read sidecar written by the paired store layer
# during prefill. Q still goes through rotary. The sidecar
# layout ([B, n_kv, max_seq, D]) does not match the paged
# cache's block-mask shape, so we route shared layers through
# the eager SDPA kernel instead of flex_attention. This is
# slower per-layer but keeps the sidecar design simple and
# CUDA-graph safe -- SDPA reads stable device pointers, and
# we attend over the full sidecar length (zero-initialized
# beyond the prefill sequence, giving negligible contribution
# once masked by causal).
q = _apply_rotary_q(q, cos, sin)
shared_k, shared_v = shared_kv_buffer[self.kv_shared_layer_index]
B = q.shape[0]
k = shared_k[:B]
v = shared_v[:B]
Hq = q.shape[1]
Hkv = k.shape[1]
if Hq != Hkv:
groups = Hq // Hkv
k = k.repeat_interleave(groups, dim = 1)
v = v.repeat_interleave(groups, dim = 1)
# Causal mask over q_pos vs kv_pos. Note: Gemma-4 has two
# attention regimes (full_attention and sliding_attention),
# both causal; sliding-window layers additionally clamp kv to
# the last `sliding_window` positions. The sidecar path here
# treats shared layers as full-causal over the prefix. Strict
# sliding-window semantics on shared layers are a TODO; with
# 512-token windows and typical prefixes this approximation
# matches within a few ULP, but would drift on long prefixes.
is_causal = q.shape[-2] > 1
attn_output = F.scaled_dot_product_attention(
q,
k,
v,
is_causal = is_causal,
scale = self.scaling,
)
attn_output = (
attn_output.transpose(1, 2).reshape(*input_shape, -1).contiguous()
)
return self.o_proj(attn_output), None
else:
k_raw = self.k_proj(hidden_states).view(hidden_shape)
k = self.k_norm(k_raw).transpose(1, 2)
# `v_proj` may be None under Gemma-4's K=V global-attention
# option; in that case reuse the raw (un-normed) k projection.
if self.v_proj is not None:
v = self.v_norm(
self.v_proj(hidden_states).view(hidden_shape)
).transpose(1, 2)
else:
v = k_raw.transpose(1, 2)
q, k = _apply_rotary(q, k, cos, sin)
# Prefill-only sidecar write for store layers.
if getattr(self, "store_full_length_kv", False):
is_prefill = q.shape[-2] > 1
if is_prefill:
shared_k, shared_v = shared_kv_buffer[self.layer_idx]
B = k.shape[0]
S = k.shape[-2]
shared_k[:B, :, :S, :].copy_(k)
shared_v[:B, :, :S, :].copy_(v)
if self._paged_cache is not None and flex_input_pos is not None:
k, v = self._paged_cache.update(flex_input_pos, k, v, flex_batch_idx)
attn_output = flex_attention_compiled(
q,
k,
v,
scale = self.scaling,
block_mask = flex_block_mask,
enable_gqa = True,
kernel_options = flex_kernel_options,
)
attn_output = attn_output.transpose(1, 2).reshape(*input_shape, -1).contiguous()
return self.o_proj(attn_output), None
return forward
def patch_gemma4_attention_forwards(
model: torch.nn.Module, page_table: PageTable, shared_kv_buffer: dict
):
"""Attach a PagedKVCache to every non-shared attention layer and swap
in the flex_attention forward above. Shared layers get `_paged_cache =
None` because their K/V comes from the sidecar.
"""
fwd = make_flex_gemma4_attention_forward(page_table, shared_kv_buffer)
for layer in model.model.layers:
attn = layer.self_attn
if getattr(attn, "is_kv_shared_layer", False):
attn._paged_cache = None
else:
n_kv = attn.k_proj.out_features // attn.head_dim
attn._paged_cache = PagedKVCache(
page_table,
n_heads = n_kv,
head_dim = attn.head_dim,
dtype = model.dtype,
).to(model.device)
attn.forward = types.MethodType(fwd, attn)
# --- model forward walker --------------------------------------------------
def call_gemma4_model_with_flex_kwargs(
model, input_ids, position_ids, flex_kwargs
):
"""Walk the Gemma-4 text model manually so we can inject flex_* kwargs
into each attention call. Mirrors `call_model_with_flex_kwargs` in
`qwen3_flex_inference.py` but:
- precomputes both (cos, sin) variants and passes them as a dict
keyed by `layer.self_attn.layer_type`;
- materializes `per_layer_inputs` via the model's own
`get_per_layer_inputs` + `project_per_layer_inputs` helpers;
- runs the double-residual around attn, the double-residual around
MLP, the per-layer-input path, and the `layer_scalar` multiply.
The final norm is applied here; lm_head + softcap is applied by the
caller (so we can slice to logits_positions before the vocab matmul).
"""
base = model.model
inputs_embeds = base.embed_tokens(input_ids)
# Per-layer input table. E2B has hidden_size_per_layer_input = 256 and
# 35 layers, so this is [B, S, 35, 256] -- a local tensor with fixed
# shape across CUDA graph replays (input_ids is pre-allocated upstream).
per_layer_inputs = None
if getattr(base, "hidden_size_per_layer_input", 0):
per_layer_inputs = base.get_per_layer_inputs(input_ids, inputs_embeds)
per_layer_inputs = base.project_per_layer_inputs(
inputs_embeds, per_layer_inputs
)
position_embeddings = {
layer_type: base.rotary_emb(inputs_embeds, position_ids, layer_type)
for layer_type in base.unique_layer_types
}
hidden_states = inputs_embeds
for i, layer in enumerate(base.layers):
residual = hidden_states
hidden_states = layer.input_layernorm(hidden_states)
hidden_states, _ = layer.self_attn(
hidden_states,
position_embeddings = position_embeddings,
**flex_kwargs,
)
hidden_states = layer.post_attention_layernorm(hidden_states)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = layer.pre_feedforward_layernorm(hidden_states)
hidden_states = layer.mlp(hidden_states)
hidden_states = layer.post_feedforward_layernorm(hidden_states)
hidden_states = residual + hidden_states
if per_layer_inputs is not None and hasattr(layer, "per_layer_input_gate"):
residual = hidden_states
hidden_states = layer.per_layer_input_gate(hidden_states)
hidden_states = layer.act_fn(hidden_states)
hidden_states = hidden_states * per_layer_inputs[:, :, i, :]
hidden_states = layer.per_layer_projection(hidden_states)
hidden_states = layer.post_per_layer_input_norm(hidden_states)
hidden_states = residual + hidden_states
if hasattr(layer, "layer_scalar"):
hidden_states = hidden_states * layer.layer_scalar
hidden_states = base.norm(hidden_states)
return hidden_states
# --- inference engine ------------------------------------------------------
class FlexGemma4Inference:
def __init__(
self,
model,
tokenizer,
max_batch_size = 32,
max_seq_length = 2048,
n_pages = 2048,
page_size = 128,
max_new_tokens = 512,
decode_kernel_options = None,
prefill_kernel_options = None,
fa4_prefill = None,
base_model = None,
peft_model = None,
):
assert max_seq_length % page_size == 0
self.model = model
self.tokenizer = tokenizer
self.device = model.device
self.eos_token_id = tokenizer.eos_token_id
self.base_model = base_model
self.peft_model = peft_model
self.max_batch_size = max_batch_size
self.max_seq_length = max_seq_length
self.page_size = page_size
self.max_new_tokens = max_new_tokens
if fa4_prefill is None or fa4_prefill:
major, _ = torch.cuda.get_device_capability(self.device)
supported = major >= 9
if fa4_prefill and not supported:
import warnings
warnings.warn(
f"--fa4_prefill needs Hopper (sm_90) or Blackwell "
f"(sm_100 / sm_120); found sm_{major}0. Falling back "
f"to the Triton flex_attention backend.",
RuntimeWarning,
stacklevel = 2,
)
fa4_prefill = supported
self.fa4_prefill = fa4_prefill
self.prefill_q_block = 256 if fa4_prefill else 128
self.prefill_kv_block = 128
self.decode_kernel_options = (
decode_kernel_options
if decode_kernel_options is not None
else DECODE_KERNEL_OPTIONS_DEFAULT
)
base_prefill_opts = (
prefill_kernel_options
if prefill_kernel_options is not None
else dict(PREFILL_KERNEL_OPTIONS_DEFAULT)
)
if fa4_prefill:
base_prefill_opts = dict(base_prefill_opts)
base_prefill_opts.pop("FORCE_USE_FLEX_ATTENTION", None)
base_prefill_opts["BACKEND"] = "FLASH"
self.prefill_kernel_options = base_prefill_opts
self.page_table = PageTable(
n_pages = n_pages,
page_size = page_size,
max_batch_size = max_batch_size,
device = self.device.type,
)
# Allocate shared-KV sidecar buffers. Only "store" layers get an
# entry; shared layers read by the store layer's `layer_idx`.
# Buffers live at fixed device addresses, so CUDA graph replay
# reads stable pointers -- safe because store layers only write
# during prefill (not inside captured decode graphs) and shared
# layers only read.
#
# Per-layer num_kv_heads is derived from `k_proj.out_features /
# head_dim` rather than `config.num_key_value_heads`, because
# global-attention layers may use `num_global_key_value_heads`
# under `attention_k_eq_v`.
self.shared_kv_buffer: dict = {}
for i, layer in enumerate(model.model.layers):
attn = layer.self_attn
if not getattr(attn, "store_full_length_kv", False):
continue
hd = attn.head_dim
n_kv = attn.k_proj.out_features // hd
K = torch.zeros(
max_batch_size,
n_kv,
max_seq_length,
hd,
dtype = model.dtype,
device = self.device,
)
V = torch.zeros_like(K)
self.shared_kv_buffer[i] = (K, V)
patch_gemma4_attention_forwards(
model, self.page_table, self.shared_kv_buffer
)
self.input_pos_buffer = torch.zeros(
max_batch_size, dtype = torch.int32, device = self.device
)
self.block_mask_logical = self.page_table.create_causal_blockmask(
B = max_batch_size,
L = max_seq_length,
)
self.cudagraph_captured = False
self.graphs = {}
self.graph_vars = {}
def tokenize(self, sequences):
for seq in sequences:
ids = self.tokenizer(seq.text, return_tensors = "pt")["input_ids"].squeeze(0)
seq.input_ids = ids
seq.input_length = ids.shape[0]
def _softcap(self, logits):
sc = getattr(self.model.config, "final_logit_softcapping", None)
if sc is not None and sc > 0:
logits = torch.tanh(logits / sc) * sc
return logits
def _prefill(self, batch: list) -> torch.Tensor:
input_ids_list = [seq.input_ids.to(self.device) for seq in batch]
input_pos_list = [
torch.arange(seq.input_length, dtype = torch.long, device = self.device)
for seq in batch
]
batch_idx_list = [
torch.full(
(seq.input_length,),
seq.batch_idx,
dtype = torch.long,
device = self.device,
)
for seq in batch
]
input_ids = torch.cat(input_ids_list).view(1, -1)
input_pos = torch.cat(input_pos_list).view(1, -1)
batch_idx = torch.cat(batch_idx_list).view(1, -1)
L = input_ids.shape[1]
q_block = self.prefill_q_block
pad = (q_block - L % q_block) % q_block
if pad > 0:
input_ids = F.pad(input_ids, (0, pad), value = 0)
input_pos = F.pad(input_pos, (0, pad), value = 0)
batch_idx = F.pad(batch_idx, (0, pad), value = 0)
input_lengths = torch.tensor(
[s.input_length for s in batch], dtype = torch.long, device = self.device
)
logits_positions = input_lengths.cumsum(dim = 0) - 1
prefill_block_size = (
(self.prefill_q_block, self.prefill_kv_block)
if self.fa4_prefill
else self.prefill_q_block
)
mask = self.page_table.create_prefill_blockmask_no_paging(
batch_idx, BLOCK_SIZE = prefill_block_size
)
flex_kwargs = dict(
flex_block_mask = mask,
flex_input_pos = input_pos,
flex_batch_idx = batch_idx,
flex_kernel_options = self.prefill_kernel_options,
)
hidden = call_gemma4_model_with_flex_kwargs(
self.model, input_ids, input_pos, flex_kwargs
)
logits = self.model.lm_head(hidden[:, logits_positions, :]).squeeze(0)
return self._softcap(logits)
def _decode_block_mask(self, batch_idx: torch.Tensor):
block_mask = self.block_mask_logical
input_pos = self.input_pos_buffer[batch_idx]
assert batch_idx.ndim == 1 and input_pos.ndim == 1
B = batch_idx.shape[0]
input_block_idx = input_pos // block_mask.BLOCK_SIZE[0]
kv_num_blocks = block_mask.kv_num_blocks[batch_idx, :, input_block_idx].view(
B, 1, 1
)
kv_indices = block_mask.kv_indices[batch_idx, :, input_block_idx].view(
B, 1, 1, -1
)
full_num = full_idx = None
if block_mask.full_kv_num_blocks is not None:
full_num = block_mask.full_kv_num_blocks[
batch_idx, :, input_block_idx
].view(B, 1, 1)
full_idx = block_mask.full_kv_indices[
batch_idx, :, input_block_idx
].view(B, 1, 1, -1)
def causal_offset(off):
def offset(b, h, q_idx, kv_idx):
return q_idx + off[b] >= kv_idx
return offset
seq_length = (1, block_mask.seq_lengths[1])
mask = BlockMask.from_kv_blocks(
kv_num_blocks,
kv_indices,
full_num,
full_idx,
BLOCK_SIZE = block_mask.BLOCK_SIZE,
mask_mod = causal_offset(input_pos),
seq_lengths = seq_length,
)
return mask, input_pos
def _decode_step_eager(self, batch_idx: torch.Tensor, input_ids: torch.Tensor):
B = input_ids.shape[0]
mask, input_pos = self._decode_block_mask(batch_idx)
mask = self.page_table.convert_logical_block_mask(mask, batch_idx)
position_ids = input_pos.view(B, 1).to(torch.long)
flex_kwargs = dict(
flex_block_mask = mask,
flex_input_pos = input_pos.view(B, 1).to(torch.long),
flex_batch_idx = batch_idx,
flex_kernel_options = self.decode_kernel_options,
)
hidden = call_gemma4_model_with_flex_kwargs(
self.model, input_ids.view(B, 1), position_ids, flex_kwargs
)
logits = self.model.lm_head(hidden[:, -1, :])
return self._softcap(logits)
def _decode_step(
self,
batch_idx: torch.Tensor,
input_ids: torch.Tensor,
input_pos: torch.Tensor,
):
self.input_pos_buffer.zero_()
self.input_pos_buffer[batch_idx] = input_pos
if not self.cudagraph_captured:
return self._decode_step_eager(batch_idx, input_ids)
bs = input_ids.size(0)
key = next(x for x in self.graph_bs if x >= bs)
graph = self.graphs[key]
gv = self.graph_vars
for k, v in gv.items():
if k != "outputs":
v.zero_()
gv["input_ids"][:bs] = input_ids
gv["batch_idx"][:bs] = batch_idx
graph.replay()
return gv["outputs"][:bs]
def capture_decode_cudagraph(self):
max_bs = self.max_batch_size
reserved_batches = []
for bi in range(1, max_bs):
try:
allocated = self.page_table.allocate()
self.page_table.reserve(
allocated,
torch.tensor(
[allocated], device = self.device, dtype = torch.long
),
self.page_size,
)
reserved_batches.append(allocated)
except Exception:
break
input_ids = torch.zeros(max_bs, dtype = torch.int64, device = self.device)
batch_idx = torch.arange(max_bs, dtype = torch.int64, device = self.device)
outputs = torch.zeros(
(max_bs, self.model.config.vocab_size),
dtype = self.model.dtype,
device = self.device,
)
self.graph_bs = [1, 2, 4, 8] + list(range(16, max_bs + 1, 16))
pool = None
for bs in reversed(self.graph_bs):
if bs > max_bs:
continue
print(f"[flex-gemma4] capturing CUDA graph for bs={bs}")
torch.cuda.synchronize()
_ = self._decode_step_eager(batch_idx[:bs], input_ids[:bs])
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, pool):
outputs[:bs] = self._decode_step_eager(batch_idx[:bs], input_ids[:bs])
if pool is None:
pool = graph.pool()
self.graphs[bs] = graph
torch.cuda.synchronize()
for bi in reserved_batches:
self.page_table.erase(bi)
self.graph_vars = dict(
input_ids = input_ids, batch_idx = batch_idx, outputs = outputs
)
def refresh_inference_from_base(self):
if self.base_model is None or self.peft_model is None:
return 0
return refresh_lora_merge_from_pristine(self.base_model, self.peft_model)
@torch.inference_mode()
def generate(self, sequences, capture_cudagraph = False):
self.tokenize(sequences)
waiting = deque(sequences)
running = deque()
done = []
if capture_cudagraph and not self.cudagraph_captured:
self.capture_decode_cudagraph()
self.cudagraph_captured = True
while waiting or running:
batch = []
while waiting and self.page_table.can_reserve(waiting[0].total_length):
seq = waiting.popleft()
bi = self.page_table.allocate()
self.page_table.reserve(
bi,
torch.tensor([bi], device = self.device, dtype = torch.long),
seq.total_length,
)
seq.batch_idx = bi
batch.append(seq)
if batch:
logits = self._prefill(batch)
next_ids = torch.argmax(logits, dim = -1).tolist()
for i, seq in enumerate(batch):
seq.last_token_id = next_ids[i]
seq.output_ids.append(next_ids[i])
if (
seq.last_token_id == self.eos_token_id
or len(seq.output_ids) >= seq.max_new_tokens
):
seq.finished = True
done.append(seq)
self.page_table.erase(seq.batch_idx)
else:
running.append(seq)
continue
decode_batch = []
while running:
seq = running.popleft()
if self.page_table.capacity[seq.batch_idx] >= seq.total_length:
decode_batch.append(seq)
elif self.page_table.can_reserve(
seq.total_length, batch_idx_int = seq.batch_idx
):
self.page_table.reserve(
seq.batch_idx,
torch.tensor(
[seq.batch_idx],
device = self.device,
dtype = torch.long,
),
seq.total_length,
)
decode_batch.append(seq)
else:
running.appendleft(seq)
newest = running.pop()
waiting.appendleft(newest)
self.page_table.erase(newest.batch_idx)
if not decode_batch:
continue
B = len(decode_batch)
bi_tensor = torch.tensor(
[s.batch_idx for s in decode_batch],
dtype = torch.long,
device = self.device,
)
last_ids = torch.tensor(
[s.last_token_id for s in decode_batch],
dtype = torch.long,
device = self.device,
)
cur_pos = torch.tensor(
[s.total_length - 1 for s in decode_batch],
dtype = torch.int32,
device = self.device,
)
logits = self._decode_step(bi_tensor, last_ids, cur_pos)
next_ids = torch.argmax(logits, dim = -1).tolist()
for i, seq in enumerate(decode_batch):
seq.last_token_id = next_ids[i]
seq.output_ids.append(next_ids[i])
if (
seq.last_token_id == self.eos_token_id
or len(seq.output_ids) >= seq.max_new_tokens
):
seq.finished = True
done.append(seq)
self.page_table.erase(seq.batch_idx)
else:
running.append(seq)
return done
# --- CLI -------------------------------------------------------------------
def main():
p = argparse.ArgumentParser()
p.add_argument("--model_name", default = "unsloth/gemma-4-E2B-it")
p.add_argument("--n_prompts", type = int, default = 64)
p.add_argument("--n_rounds", type = int, default = 2)
p.add_argument("--max_new_tokens", type = int, default = 512)
p.add_argument("--max_batch_size", type = int, default = 64)
p.add_argument("--max_seq_length", type = int, default = 2048)
p.add_argument("--n_pages", type = int, default = 2048)
p.add_argument("--page_size", type = int, default = 128)
p.add_argument("--capture_cudagraph", action = "store_true")
p.add_argument("--lora_adapter", default = None)
p.add_argument("--decode_kernel_options", default = None)
p.add_argument("--prefill_kernel_options", default = None)
p.add_argument(
"--fa4_prefill",
default = None,
action = argparse.BooleanOptionalAction,
help = (
"Use BLOCK_SIZE=(256,128) + BACKEND=FLASH on prefill. Default "
"auto-enables on Hopper (sm_90) and Blackwell (sm_100, sm_120)."
),
)
p.add_argument("--load_in_4bit", action = "store_true")
p.add_argument(
"--no_merge_lora",
action = "store_true",
help = "Keep the LoRA adapter as a PEFT wrapper instead of merging.",
)
p.add_argument(
"--verify_no_drift",
action = "store_true",
help = "Drift-verify the double-copy LoRA refresh across N cycles.",
)
p.add_argument("--verify_iterations", type = int, default = 10)
p.add_argument("--model_name_4bit", default = None)
p.add_argument("--stats_path", required = True)
p.add_argument(
"--chat_template",
choices = ["auto", "grpo", "native"],
default = "auto",
help = (
"Which chat template to use. `auto`: native for Gemma-4. "
"`grpo`: force the GRPO template. `native`: force the "
"tokenizer's built-in template."
),
)
args = p.parse_args()
def _parse_opts(s):
if s is None:
return None
return json.loads(s)
Gemma4ForCausalLM, Gemma4Config, Gemma4TextConfig = _require_gemma4()
from transformers import AutoTokenizer
from transformers.models.gemma4.modeling_gemma4 import (
Gemma4ForConditionalGeneration,
)
tok = AutoTokenizer.from_pretrained(args.model_name)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
base_model = None
peft_model = None
if args.load_in_4bit:
from transformers import AutoModelForCausalLM
from huggingface_hub import HfApi
bnb_model_name = args.model_name_4bit or f"{args.model_name}-unsloth-bnb-4bit"
# Probe for the 4-bit shard. If missing, the user asked for a
# quant row we cannot produce; fail loudly rather than silently
# falling back to bf16 (which would mislabel the stats file).
try:
HfApi().model_info(bnb_model_name)
except Exception as e:
raise SystemExit(
f"[flex-gemma4] --load_in_4bit: 4-bit shard {bnb_model_name} "
f"is not available ({e}). Use --model_name_4bit to override "
f"or drop --load_in_4bit for bf16."
)
print(f"[flex-gemma4] loading 4-bit base: {bnb_model_name}")
model = AutoModelForCausalLM.from_pretrained(
bnb_model_name,
attn_implementation = "eager",
device_map = "cuda:0",
)
if getattr(model.config, "tie_word_embeddings", False):
model.lm_head.weight = model.model.embed_tokens.weight
model.eval()
if args.lora_adapter:
from peft import PeftModel
peft_wrapper = PeftModel.from_pretrained(
model,
str(Path(args.lora_adapter).resolve()),
is_trainable = False,
)
model = peft_wrapper.base_model.model
else:
# Text-only bf16. The checkpoint stores text weights under the
# `model.language_model.` prefix (Gemma-4 is natively multimodal).
# We load the full `Gemma4ForConditionalGeneration`, pluck the
# language_model, then drop the vision + audio towers before
# moving to GPU so peak memory reflects text only.
full_cfg = Gemma4Config.from_pretrained(args.model_name)
text_cfg = full_cfg.text_config
full_model = Gemma4ForConditionalGeneration.from_pretrained(
args.model_name,
dtype = torch.bfloat16,
attn_implementation = "eager",
)
lang_model = full_model.model.language_model
# Drop the non-text towers. `embed_vision` / `embed_audio` project
# from text hidden size -- harmless when not invoked, but we kill
# them too so deepcopy (below) stays cheap.
full_model.model.vision_tower = None
full_model.model.audio_tower = None
full_model.model.embed_vision = None
full_model.model.embed_audio = None
# Build a ForCausalLM shell around the language_model so LoRA /
# state-dict hashing treat it like any other HF decoder model.
base_model = Gemma4ForCausalLM(text_cfg)
base_model.model = lang_model
base_model.lm_head.weight = lang_model.embed_tokens.weight
base_model = base_model.to(torch.bfloat16).to("cuda")
base_model.eval()
del full_model
if not args.lora_adapter:
model = base_model
base_model = None
elif args.no_merge_lora:
from peft import PeftModel
peft_wrapper = PeftModel.from_pretrained(
base_model,
str(Path(args.lora_adapter).resolve()),
is_trainable = False,
)
model = peft_wrapper.base_model.model
base_model = None
else:
from peft import PeftModel
print("[flex-gemma4] deep-copying base model for double-copy rollout")
inference_model = copy.deepcopy(base_model)
inference_model.eval()
peft_model = PeftModel.from_pretrained(
inference_model,
str(Path(args.lora_adapter).resolve()),
is_trainable = False,
)
model = peft_model.base_model.model
model.eval()
if args.verify_no_drift:
if args.load_in_4bit:
raise SystemExit(
"--verify_no_drift only applies to the bf16 double-copy path."
)
if args.no_merge_lora:
raise SystemExit(
"--verify_no_drift is incompatible with --no_merge_lora."
)
if base_model is None or peft_model is None:
raise SystemExit(
"--verify_no_drift requires --lora_adapter against the bf16 path."
)
print(
f"[flex-gemma4] running drift verification: "
f"{args.verify_iterations} perturb+refresh cycles"
)
result = run_drift_verification(
base_model, peft_model, n_iters = args.verify_iterations
)
result = {"mode": "verify_no_drift", "model_name": args.model_name, **result}
os.makedirs(
os.path.dirname(os.path.abspath(args.stats_path)) or ".",
exist_ok = True,
)
with open(args.stats_path, "w") as f:
json.dump(result, f, indent = 2)
print(json.dumps(result, indent = 2))
os._exit(0)
from unsloth_grpo_common import SYSTEM_PROMPT, apply_chat_template_to_tokenizer
from datasets import load_dataset
if args.chat_template == "auto":
# Gemma-4 default is the tokenizer native template; only Qwen3
# in this repo uses GRPO-by-default.
use_grpo = False
elif args.chat_template == "grpo":
use_grpo = True
else:
use_grpo = False
if use_grpo:
apply_chat_template_to_tokenizer(tok)
print("[flex-gemma4] chat_template: GRPO")
else:
print("[flex-gemma4] chat_template: tokenizer native")
ds = load_dataset("open-r1/DAPO-Math-17k-Processed", "en", split = "train")
ds = ds.shuffle(seed = 3407).select(range(args.n_prompts))
messages = [
[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": x["prompt"]},
]
for x in ds
]
texts = [
tok.apply_chat_template(m, add_generation_prompt = True, tokenize = False)
for m in messages
]
inference = FlexGemma4Inference(
model,
tok,
max_batch_size = args.max_batch_size,
max_seq_length = args.max_seq_length,
n_pages = args.n_pages,
page_size = args.page_size,
max_new_tokens = args.max_new_tokens,
decode_kernel_options = _parse_opts(args.decode_kernel_options),
prefill_kernel_options = _parse_opts(args.prefill_kernel_options),
fa4_prefill = args.fa4_prefill,
base_model = base_model,
peft_model = peft_model,
)
if inference.base_model is not None and inference.peft_model is not None:
n = inference.refresh_inference_from_base()
print(f"[flex-gemma4] double-copy rollout: refreshed {n} LoRA-target layers")
def make_seqs():
return [Sequence(text = t, max_new_tokens = args.max_new_tokens) for t in texts]
torch.cuda.reset_peak_memory_stats()
print("[flex-gemma4] warmup (16 prompts)...")
_ = inference.generate(make_seqs()[:16], capture_cudagraph = args.capture_cudagraph)
torch.cuda.synchronize()
wall_times = []
total_decoded = 0
for r in range(args.n_rounds):
torch.cuda.synchronize()
t0 = time.perf_counter()
out = inference.generate(make_seqs())
torch.cuda.synchronize()
wall_times.append(time.perf_counter() - t0)
total_decoded = sum(len(s.output_ids) for s in out)
print(
f"[flex-gemma4] round {r}: {wall_times[-1]:.2f}s, {total_decoded} "
f"tokens, {total_decoded / wall_times[-1]:.1f} tok/s"
)
med = sorted(wall_times)[len(wall_times) // 2]
best = min(wall_times)
peak = torch.cuda.max_memory_allocated() / 1024**3
sample_completions = []
for s in out[:3]:
sample_completions.append(
tok.decode(s.output_ids[:80], skip_special_tokens = True)
)
res = {
"backend": "flex-gemma4",
"model_name": args.model_name,
"capture_cudagraph": args.capture_cudagraph,
"lora_adapter": args.lora_adapter,
"n_prompts": args.n_prompts,
"n_decoded_tokens": total_decoded,
"wall_times_s": wall_times,
"median_wall_s": med,
"best_wall_s": best,
"decode_tps_median": total_decoded / med if med else 0,
"decode_tps_best": total_decoded / best if best else 0,
"max_new_tokens": args.max_new_tokens,
"peak_memory_gb": peak,
"sample_completions": sample_completions,
}
os.makedirs(
os.path.dirname(os.path.abspath(args.stats_path)) or ".", exist_ok = True
)
with open(args.stats_path, "w") as f:
json.dump(res, f, indent = 2)
print(json.dumps(res, indent = 2))
os._exit(0)
if __name__ == "__main__":
main()