One-shot correctness checks for the flex_attention + paged KV path against a vanilla HF `model(input_ids, use_cache=False)` forward on the same prompt. Report max / mean abs diff on last-position logits plus argmax match and top-10 overlap, so numerical drift and semantic equivalence are both visible. Shared approach: load the base model, deep-copy it for the flex path so the attention patching does not mutate the vanilla comparison, run both on the same tokenized prompt, report diffs. `verify_gemma4_numerics.py` mirrors `gemma4_flex_inference.py`'s text-only loader (Gemma4ForConditionalGeneration -> drop vision / audio towers -> language_model into a Gemma4ForCausalLM shell) before deep-copying. The softcap is NOT re-applied on the vanilla logits since `Gemma4ForCausalLM.forward` already applies `final_logit_softcapping`. `verify_qwen3_numerics.py` runs `qwen3_flex_inference.FlexInference` against either Qwen3 or Llama-3.2 via `--model_name`. `fa4_prefill` is disabled because short prompts hit a CuteDSL sm_100 shape mismatch in `handle_block_sparse_empty_tile_correction_sm100`. Results on B200 bf16, 6 to 7 token prompt: | Model | max abs | mean abs | argmax | top-10 | |---------------------------------|---------|----------|--------|--------| | unsloth/Qwen3-4B-Base | 0.313 | 0.088 | yes | 10/10 | | unsloth/Llama-3.2-3B-Instruct | 0.125 | 0.021 | yes | 10/10 | | unsloth/gemma-4-E2B-it | 0.375 | 0.080 | yes | 10/10 | All three land in the same bf16 Triton-flex vs eager-matmul drift band; Gemma-4's extra per-layer-input path and layer_scalar do not widen the gap despite 20 of 35 layers going through the shared-KV link. |
||
|---|---|---|
| .. | ||
| benchmarks | ||
| enforce_kwargs_spacing.py | ||
| install_gemma4_mlx.sh | ||
| install_qwen3_6_mlx.sh | ||
| run_ruff_format.py | ||