From b294fbd3dc214864cfe8645bf45d302ae5051fab Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 21 Apr 2026 10:00:48 +0000 Subject: [PATCH] benchmarks: verify_gemma4_numerics -- compare against raw HF, not just shell Previously the vanilla reference was a Gemma4ForCausalLM shell wrapping the language_model (the same construction used inside gemma4_flex_inference.main for LoRA / state-dict hashing convenience). That is not plain HF: `Gemma4Model.forward` uses `create_masks_for_generate(..., mm_token_type_ids, pixel_values)` to build attention masks, while the shell calls `Gemma4TextModel.forward` directly, which builds its own per-regime masks via `create_causal_mask` + `create_sliding_window_causal_mask`. Both are correct for text-only input but their mask-bias precision differs enough to produce a measurable drift. The script now keeps both references alive and reports three diffs: shell vs raw, flex vs raw, flex vs shell. On unsloth/gemma-4-E2B-it bf16 with a 6-token prompt: shell vs raw max 6.9e-01 mean 3.0e-01 argmax=yes top-10=10/10 flex vs raw max 6.3e-01 mean 2.2e-01 argmax=yes top-10=10/10 flex vs shell max 3.8e-01 mean 8.0e-02 argmax=yes top-10=10/10 Flex is actually closer to raw HF than the shell is. About 0.30 mean of the flex-vs-shell-and-vs-raw gap comes from the shell's mask construction alone, not the flex kernel. Either way the bf16 drift band matches Qwen3 (0.3 / 0.09) and Llama-3.2 (0.13 / 0.02), and semantic top-1 + top-10 are exact. --- scripts/benchmarks/verify_gemma4_numerics.py | 87 +++++++++++++------- 1 file changed, 59 insertions(+), 28 deletions(-) diff --git a/scripts/benchmarks/verify_gemma4_numerics.py b/scripts/benchmarks/verify_gemma4_numerics.py index 850753cc45..74b975338f 100644 --- a/scripts/benchmarks/verify_gemma4_numerics.py +++ b/scripts/benchmarks/verify_gemma4_numerics.py @@ -34,10 +34,24 @@ def main(): if tok.pad_token is None: tok.pad_token = tok.eos_token - # --- load once (text shell) and keep a pristine copy for the vanilla pass full_cfg = Gemma4Config.from_pretrained(name) text_cfg = full_cfg.text_config + # Untouched HF reference: `Gemma4ForConditionalGeneration.forward` -- + # the same class everyone else would load via AutoModelForCausalLM + # for Gemma-4. No patching, no shell, no flex attention. + ref_raw = Gemma4ForConditionalGeneration.from_pretrained( + name, dtype = torch.bfloat16, attn_implementation = "eager" + ).to("cuda") + ref_raw.eval() + + # Shell copy used by `gemma4_flex_inference.main()`: we deep-copy the + # loaded multimodal model, drop the vision + audio towers, and move + # the language_model into a Gemma4ForCausalLM wrapper so PEFT and + # state-dict hashing treat it as a decoder-only model. The flex path + # then patches its attention forwards on this shell. We keep the + # shell around both as (a) the model that gets flex-patched and + # (b) a sanity check that the shell itself matches the raw HF path. full = Gemma4ForConditionalGeneration.from_pretrained( name, dtype = torch.bfloat16, attn_implementation = "eager" ) @@ -47,29 +61,39 @@ def main(): full.model.embed_vision = None full.model.embed_audio = None - base = Gemma4ForCausalLM(text_cfg) - base.model = lang - base.lm_head.weight = lang.embed_tokens.weight - base = base.to(torch.bfloat16).to("cuda") - base.eval() + shell = Gemma4ForCausalLM(text_cfg) + shell.model = lang + shell.lm_head.weight = lang.embed_tokens.weight + shell = shell.to(torch.bfloat16).to("cuda") + shell.eval() del full - # Deep-copy so Flex's attention patching doesn't mutate the vanilla path. - flex_model = copy.deepcopy(base) + # Deep-copy so Flex's attention patching doesn't mutate the shell. + flex_model = copy.deepcopy(shell) prompt = "The quick brown fox jumps over" ids = tok(prompt, return_tensors = "pt")["input_ids"].to("cuda") print(f"prompt len = {ids.shape[1]}") with torch.inference_mode(): - # `Gemma4ForCausalLM.forward` applies `final_logit_softcapping` - # internally, so these logits are already softcapped. - out = base(ids, use_cache = False) - ref_logits = out.logits[0, -1, :].float() - print(f"vanilla last-token logits: mean {ref_logits.mean().item():.4f}, " - f"std {ref_logits.std().item():.4f}, " - f"argmax {int(ref_logits.argmax())} " - f"({tok.decode([int(ref_logits.argmax())])!r})") + # `Gemma4ForConditionalGeneration.forward` applies + # `final_logit_softcapping` internally. + ref_logits = ref_raw(input_ids = ids, use_cache = False).logits[ + 0, -1, : + ].float() + shell_logits = shell(ids, use_cache = False).logits[0, -1, :].float() + print( + f"raw Gemma4ForConditionalGeneration: mean {ref_logits.mean():.4f}, " + f"std {ref_logits.std():.4f}, argmax {int(ref_logits.argmax())} " + f"({tok.decode([int(ref_logits.argmax())])!r})" + ) + print( + f"shell Gemma4ForCausalLM(text_cfg) : mean {shell_logits.mean():.4f}, " + f"std {shell_logits.std():.4f}, argmax {int(shell_logits.argmax())}" + ) + # Dispose of the raw multimodal model before we build FlexGemma4Inference. + del ref_raw + torch.cuda.empty_cache() # Flex path. inf = FlexGemma4Inference( @@ -99,19 +123,26 @@ def main(): seq.batch_idx = bi with torch.inference_mode(): flex_logits = inf._prefill([seq])[0].float() - print(f"flex last-token logits: mean {flex_logits.mean().item():.4f}, " - f"std {flex_logits.std().item():.4f}, " - f"argmax {int(flex_logits.argmax())} " - f"({tok.decode([int(flex_logits.argmax())])!r})") + print( + f"flex last-token logits : mean {flex_logits.mean():.4f}, " + f"std {flex_logits.std():.4f}, argmax {int(flex_logits.argmax())} " + f"({tok.decode([int(flex_logits.argmax())])!r})" + ) - diff = (flex_logits - ref_logits).abs() - print(f"max abs diff = {diff.max().item():.4e}") - print(f"mean abs diff = {diff.mean().item():.4e}") - print(f"argmax match = {int(ref_logits.argmax()) == int(flex_logits.argmax())}") - # bf16 ULP is ~1e-2 at magnitude ~5. Report top-10 overlap too. - top_ref = set(ref_logits.topk(10).indices.tolist()) - top_flex = set(flex_logits.topk(10).indices.tolist()) - print(f"top-10 overlap = {len(top_ref & top_flex)} / 10") + def report(tag, a, b): + diff = (a - b).abs() + top_a = set(a.topk(10).indices.tolist()) + top_b = set(b.topk(10).indices.tolist()) + print( + f" {tag:14s} max {diff.max().item():.3e} mean {diff.mean().item():.3e} " + f"argmax={int(a.argmax()) == int(b.argmax())} top-10={len(top_a & top_b)}/10" + ) + + print("vs raw Gemma4ForConditionalGeneration:") + report("shell vs raw", shell_logits, ref_logits) + report("flex vs raw", flex_logits, ref_logits) + print("vs shell (Gemma4ForCausalLM wrapper):") + report("flex vs shell", flex_logits, shell_logits) if __name__ == "__main__":