From 5bfefc23773ef3800de2351d96d06c7fb6e8afa0 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 21 Apr 2026 06:51:51 +0000 Subject: [PATCH] flex: generalize qwen3_flex_inference.py to Llama-3.2 The flex_attention + paged KV + CUDA graphs inference engine was Qwen3-specific in a handful of places, but the underlying engine (PageTable, PagedKVCache, manual forward walker, decode graph capture, double-copy LoRA rollout, FA4 capability guard) reads only attributes that LlamaAttention / LlamaModel also expose. This change makes the engine run on both Qwen3 and Llama-3.2-3B-Instruct. Attention forward factory: - make_flex_qwen3_attention_forward -> make_flex_attention_forward - Guard the per-head QK RMSNorm call behind hasattr(self, "q_norm"). Qwen3 has it, Llama does not. The Qwen3 path is byte-equivalent to before: RMSNorm on [B, S, H, D] (per-head) then transpose. - patch_qwen3_model -> patch_model_attention_forwards. Chat template selection: - New --chat_template {auto,grpo,native}. auto picks GRPO for Qwen3 and the tokenizer's shipped template otherwise. grpo forces GRPO (matches prior Qwen3 baselines). native forces the tokenizer's own template (Llama-3.2-Instruct only produces coherent completions with its shipped Instruct template). Stats JSON: - backend: "qwen3_flex" -> "flex"; adds "model_name" so multi-arch runs land in a single schema. README: one paragraph noting Llama-3.2 support + the --chat_template native flag. Measured on B200 (sm_100), n_prompts 64, max_new_tokens 512, 5 rounds, --capture_cudagraph, double-copy LoRA rank 32: Qwen3-4B-Base bf16 3975 tok/s 44.2 GB Qwen3-4B-Base bf16 + LoRA 3656 tok/s 52.0 GB Qwen3-4B-Base 4bit + LoRA 1734 tok/s 40.6 GB Llama-3.2-3B-Inst bf16 4216 tok/s 34.7 GB Llama-3.2-3B-Inst bf16+L 4205 tok/s 40.9 GB Llama-3.2-3B-Inst 4bit+L 1892 tok/s 31.7 GB --verify_no_drift passes on both arches (base bit-identical across 10 perturb+refresh cycles, inference hash deterministic). Llama runs the Triton flex_attention backend instead of FA4: flash-attn-4 b9's sm_100 kernel raises a NoneType in handle_block_sparse_empty_tile_correction_sm100 on Llama-3.2's head shapes. Qwen3 is unaffected. Pass --no-fa4_prefill on Llama; auto-FA4 still enables on Qwen3. --- scripts/benchmarks/README.md | 7 ++ scripts/benchmarks/qwen3_flex_inference.py | 97 ++++++++++++++++------ 2 files changed, 79 insertions(+), 25 deletions(-) diff --git a/scripts/benchmarks/README.md b/scripts/benchmarks/README.md index c9cd60b671..3a545a8720 100644 --- a/scripts/benchmarks/README.md +++ b/scripts/benchmarks/README.md @@ -89,6 +89,13 @@ CUDA 12 (H100 boxes still on cu12): Pin `flash-attn-4==4.0.0b9` to match this benchmark. The `[cu13]` extra pulls in `nvidia-cutlass-dsl` built for CUDA 13. +`qwen3_flex_inference.py` runs on both Qwen3 and Llama-3.2 (the only +arch-specific branch is Qwen3's per-head QK RMSNorm; the rest of the +flex_attention + paged KV + CUDA graphs stack is identical). Pass +`--model_name unsloth/Llama-3.2-3B-Instruct` to target Llama, along with +`--chat_template native` to use Llama's shipped Instruct template instead +of the Qwen3 GRPO template. + | GPU | arch | sm | Auto FA4 | Triton flex_attention | |--------------|-----------|-------|----------|------------------------| | A100 | Ampere | sm_80 | off (uses Triton) | Works | diff --git a/scripts/benchmarks/qwen3_flex_inference.py b/scripts/benchmarks/qwen3_flex_inference.py index fef5d1c1dc..eb3dfef9e1 100644 --- a/scripts/benchmarks/qwen3_flex_inference.py +++ b/scripts/benchmarks/qwen3_flex_inference.py @@ -1,4 +1,4 @@ -"""Qwen3 inference with flex_attention + paged KV cache + CUDA graphs. +"""Llama / Qwen3 inference with flex_attention + paged KV cache + CUDA graphs. The transformers continuous-batching path tops out at ~10% of vLLM on this workload because `_generation_step` is Python-heavy (scheduler + paged @@ -16,10 +16,13 @@ by building paged attention on top of `torch.nn.attention.flex_attention`: the nearest bucket on each decode step and pad with batch_idx=0 (reserved as a no-op slot). -This file adapts that architecture to Qwen3-4B. The attention forward is -monkey-patched to use our PagedKVCache, and the inference loop runs -prefill + decode on the main thread (no background worker, graph replay -works end-to-end). +This file runs the architecture on Qwen3 and Llama-3.2. The attention +forward is monkey-patched to use our PagedKVCache, and the inference loop +runs prefill + decode on the main thread (no background worker, graph +replay works end-to-end). The only arch-specific branch is a per-head QK +RMSNorm that Qwen3 has and Llama does not; everything else (q/k/v/o proj, +head_dim, scaling, rotary_emb, embed_tokens, layers, final norm) is +identical attribute-for-attribute across the two families. LoRA: the bf16 path uses a **double-copy rollout pattern** when `--lora_adapter` is set. A pristine `base_model` lives on GPU alongside a @@ -35,6 +38,10 @@ Run: CUDA_VISIBLE_DEVICES=6 python scripts/benchmarks/qwen3_flex_inference.py \ --n_prompts 32 --max_new_tokens 512 --stats_path logs/qwen3_flex.json + CUDA_VISIBLE_DEVICES=7 python scripts/benchmarks/qwen3_flex_inference.py \ + --model_name unsloth/Llama-3.2-3B-Instruct --chat_template native \ + --n_prompts 32 --max_new_tokens 512 --stats_path logs/llama32_flex.json + Add `--capture_cudagraph` to capture per-batch-size decode graphs during warmup. """ @@ -92,11 +99,16 @@ def _apply_rotary(q, k, cos, sin): return q, k -def make_flex_qwen3_attention_forward(page_table: PageTable): - """Return a new `forward` method for `Qwen3Attention` that uses - flex_attention against a paged KV cache. The returned closure captures - the shared PageTable; each layer gets its own PagedKVCache attached to - the module as `self._paged_cache`. +def make_flex_attention_forward(page_table: PageTable): + """Return a new `forward` method for a decoder-only attention layer + (Qwen3Attention or LlamaAttention) that uses flex_attention against a + paged KV cache. The returned closure captures the shared PageTable; + each layer gets its own PagedKVCache attached to the module as + `self._paged_cache`. + + The only arch-specific branch is Qwen3's per-head QK RMSNorm + (`self.q_norm` / `self.k_norm`), applied after proj+reshape but + before rotary. Llama has no QK-norm so the guard skips. Expects the caller to have set on each layer: self._paged_cache: PagedKVCache @@ -123,8 +135,14 @@ def make_flex_qwen3_attention_forward(page_table: PageTable): input_shape = hidden_states.shape[:-1] hidden_shape = (*input_shape, -1, self.head_dim) - q = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2) - k = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) + if hasattr(self, "q_norm"): + # Qwen3: RMSNorm on [B, S, H, D] (per-head), then transpose. + q = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2) + k = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) + else: + # Llama: no QK-norm. + q = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2) + k = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2) v = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) cos, sin = position_embeddings @@ -153,11 +171,11 @@ def make_flex_qwen3_attention_forward(page_table: PageTable): return forward -def patch_qwen3_model(model: torch.nn.Module, page_table: PageTable): - """Attach a `PagedKVCache` to every `Qwen3Attention` layer and swap in - the flex_attention forward above. +def patch_model_attention_forwards(model: torch.nn.Module, page_table: PageTable): + """Attach a `PagedKVCache` to every attention layer of a Qwen3 or Llama + HF decoder model and swap in the flex_attention forward above. """ - fwd = make_flex_qwen3_attention_forward(page_table) + fwd = make_flex_attention_forward(page_table) for layer in model.model.layers: attn = layer.self_attn attn._paged_cache = PagedKVCache( @@ -176,10 +194,11 @@ def patch_qwen3_model(model: torch.nn.Module, page_table: PageTable): def call_model_with_flex_kwargs(model, input_ids, position_ids, flex_kwargs): - """`model(**inputs, **flex_kwargs)` would error because Qwen3ForCausalLM - doesn't declare the flex_* kwargs. We walk through the model manually - to pass them into the attention layers (which now accept them).""" - base = model.model # Qwen3Model + """`model(**inputs, **flex_kwargs)` would error because the HF ForCausalLM + class doesn't declare the flex_* kwargs. We walk through the model + manually to pass them into the attention layers (which now accept them). + Works identically for Qwen3 and Llama-3.2.""" + base = model.model # Qwen3Model or LlamaModel inputs_embeds = base.embed_tokens(input_ids) position_embeddings = base.rotary_emb(inputs_embeds, position_ids) hidden_states = inputs_embeds @@ -525,7 +544,7 @@ class FlexInference: max_batch_size = max_batch_size, device = self.device.type, ) - patch_qwen3_model(model, self.page_table) + patch_model_attention_forwards(model, self.page_table) # Pre-allocated decode state. self.input_pos_buffer = torch.zeros( @@ -938,6 +957,18 @@ def main(): ), ) p.add_argument("--stats_path", required = True) + p.add_argument( + "--chat_template", + choices = ["auto", "grpo", "native"], + default = "auto", + help = ( + "Which chat template to use for building prompts. " + "`auto`: GRPO template for Qwen3, tokenizer's native template " + "otherwise. `grpo`: force the GRPO template (matches prior " + "Qwen3 baselines). `native`: force the tokenizer's built-in " + "template (required for Llama-3.2-Instruct)." + ), + ) args = p.parse_args() def _parse_opts(s): @@ -1078,7 +1109,21 @@ def main(): ) from datasets import load_dataset - apply_chat_template_to_tokenizer(tok) + # Pick which chat template builds the prompts. Qwen3 baselines in + # this repo were recorded against the GRPO template; Llama-3.2-Instruct + # only produces coherent completions with its shipped Instruct + # template. + if args.chat_template == "auto": + use_grpo = type(model).__name__.startswith("Qwen3") + elif args.chat_template == "grpo": + use_grpo = True + else: # "native" + use_grpo = False + if use_grpo: + apply_chat_template_to_tokenizer(tok) + print("[flex] chat_template: GRPO") + else: + print("[flex] 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 = [ @@ -1093,8 +1138,9 @@ def main(): for m in messages ] - # Make sure the base HF model that Qwen3Attention belongs to isn't wrapped - # by PeftModel anymore (we merged); `.model` should be Qwen3ForCausalLM. + # Make sure the base HF model the attention layers belong to isn't + # wrapped by PeftModel anymore (we merged); `.model` should be + # Qwen3ForCausalLM or LlamaForCausalLM. inference = FlexInference( model, tok, @@ -1173,7 +1219,8 @@ def main(): tok.decode(s.output_ids[:80], skip_special_tokens = True) ) res = { - "backend": "qwen3_flex", + "backend": "flex", + "model_name": args.model_name, "capture_cudagraph": args.capture_cudagraph, "lora_adapter": args.lora_adapter, "n_prompts": args.n_prompts,