Repository navigation
feat(models): add Kyutai Helium (helium) text model support - #930
Conversation
Adds Kyutai Helium as a fully integrated text model: config parse and validation, loader, detection, registration, LoadedModel dispatch, docs and unit tests. Helium is a dense Llama-shaped decoder (RMSNorm pre-attention and pre-MLP, grouped-query attention, SwiGLU MLP, no QK-norm, no MoE, no sliding window, Llama's own weight key names), so the decoder block, attention and MLP come from src/models/llama3.rs unchanged rather than being copied. The one architectural difference is the RoPE convention. Upstream builds nn.RoPE(head_dim, traditional=True, base=rope_theta), so Helium rotates interleaved channel pairs (2i, 2i+1) where every other Llama-family model in this tree rotates split-half pairs (i, i + dims/2). The two produce identically shaped tensors from identical weights, so running the wrong one is a silent quality regression that no shape assertion can catch. Three code paths in llama3::Attention can rotate Q and K, and only one of them can express the convention. fused_causal_prefill_attention and FusedQKVLinear::forward_split_rope apply RoPE inside C++ launchers that hardcode traditional = false and take no flag, and both are reachable only for quantized weights, which is exactly what the Helium validation checkpoint is. Both are therefore gated on !rope_traditional so a traditional-RoPE model always takes the graph fallback, which receives the flag; the batched route (fast_rope_batched) receives it too. Neither fused path is enabled by default, so this costs the existing families nothing. llama3::ModelArgs::rope_traditional is #[serde(skip)] on purpose: a checkpoint author must not be able to change how an existing Llama checkpoint decodes by adding a config key, so the flag is only ever set programmatically by the Helium loader. Every pre-existing family keeps rope_traditional = false and is behaviorally unchanged; the tensor-parallel llama and qwen2 full-model parity tests still pass. Tensor parallelism is refused for Helium. The TP runtime builds its per-rank model by parsing config.json straight into llama3::ModelArgs, which cannot carry the flag, so a sharded Helium would silently rotate split-half while the single-process path rotates interleaved. The arch-string arm exists only to keep the dispatch table total; runtime_kind_for has no arm for ModelType::Helium and so refuses. Untrusted-config hardening is built in rather than retrofitted. ModelArgs::validate rejects zero heads before any divisibility check (0.is_multiple_of(0) is true), out-of-range hidden_size / num_hidden_layers / intermediate_size / vocab_size / max_position_embeddings, an indivisible head split, a GQA split where num_attention_heads is not divisible by num_key_value_heads, a declared head_dim that disagrees with hidden_size / num_attention_heads (upstream never reads the field, so the two must agree rather than one being silently preferred), an odd or non-positive head width (MLX enforces that contract on rope's dims argument by throwing, and an MLX C++ exception crossing the cxx bridge is an uncatchable std::terminate at the first forward pass rather than a load error), a non-positive or non-finite rope_theta, a non-finite, negative or zero rms_norm_eps (Helium's own 1e-08 is unusually small for this family and stays accepted), and a quantization block whose bits falls outside 1..=32 or whose group_size is non-positive. validate_weights runs before Llama3Model::from_weights, because FusedQKVLinear::from_weights_separate concatenates q/k/v with no shape check of its own and Attention::forward then reshapes using config-derived head counts; both throw inside MLX. It rejects an embedding table or output head with fewer rows than vocab_size claims (MLX wraps a negative gather index but does not range-check a positive one), any projection, bias or norm whose shape disagrees with the config in any layer including on the quantized path where packing compresses the input axis only, and an attention block whose q/k/v disagree on quantization or on carrying affine biases, since the fused loader decides from q_proj alone and silently drops the whole bias set when one is missing. Verified by mutation, not just by passing. Dropping the flag on the single-sequence fast_rope call failed exactly helium_logits_differ_from_the_same_weights_with_split_half_rope and every_rope_route_honors_the_flag_consistently with all 19 other tests green; dropping it on fast_rope_batched alone failed exactly every_rope_route_honors_the_flag_consistently with 20 green. Both mutations were reverted and the suite re-run green. the_fused_qkv_rope_launcher_cannot_express_traditional_rope asserts the C++ launcher's rotation against both conventions on real quantized weights, so the claim the bypass rests on is pinned rather than assumed. Closes #837
|
Real-checkpoint validation on Apple M1 Ultra (Metal + Accelerate, release build). Token-exact against the mlx-lm reference for 40 greedy tokens on The checkpoint being 4-bit affine is what makes this gate meaningful rather than incidental. Traditional RoPE is the single architectural delta in this port, and the quantized path is precisely where it was at risk: Worth noting for the reviewers: the PR bypasses the fused launchers when |
Two review follow-ups on the Helium port, neither of which changes runtime behavior. `Attention::forward` and `Attention::from_weights` are shared functions this PR modified (the two `!self.rope_traditional` gates live in the former, the flag wiring in the latter), and neither carried the `// Used by:` comment `docs/code-guidelines.md` requires. That comment is the project's discovery mechanism for "what breaks if I change this", and all three of its update triggers fired here: an attention implementation, a behavior change, and a new model reusing it. Both now list the full set that reaches this attention, which is wider than the port assumed: Llama and Mistral checkpoints, Qwen2 / Qwen2.5 through the re-export, Helium, the mllama text decoder's self-attention layers, the seven VLM text backbones built on `Llama3Model` / `Qwen2Model`, the llama and mistral pipeline stage executors, and the tensor-parallel Llama runtime. `forward_split_attention`'s list was updated by the port but named only three of those, so it now defers to the complete one rather than restating a partial copy. `helium_logits_differ_from_the_same_weights_with_split_half_rope` compared the traditional-vs-split-half logits gap against a bare `1e-3`. Measured on the pinned inputs the gap is 1.19e-2 against a logits scale of 6.89, so the absolute floor has ~12x margin today, but it is not anchored to the magnitude it is measured on: a future change that widens `tiny_args` or rescales `filled` can grow the logits without growing the separation and the test would keep passing while drifting back toward the blindness an earlier filler already caused once. Added a second assertion requiring the gap to exceed `logits_scale * 1e-4` (~17x margin), plus both measured values in the failure messages and a comment pinning why both floors exist. Verified by mutation: raising the relative floor to 1e-2 fails with "gap 0.0118608475, logits scale 6.8892975". Validated with cargo fmt, clippy (only the two warnings pre-existing on main), models::helium 22 passed, models::detection 25 passed, model_metadata 8 passed, distributed::tensor_parallel 338 passed. The one failure in a full models:: sweep, models::mllama::text::tests::ragged_real_tile_rows_match_reference_masked_full_rows, was confirmed pre-existing by running it on main in a separate worktree, where it fails with the identical 3.7252903e-9. Refs #837
Implementation Review SummaryIntent
Findings Addressed
Remaining Items
Verified independently
Not this PR
Verification
|
…he Helium loader `validate_weights` checked a quantized projection's row count but never its input width, and the q/k/v cross-check only proved the three agreed with each other. Packing compresses the input axis only, so a checkpoint packed for a different `hidden_size` keeps exactly the right number of rows and passed every check. MLX reconstructs a quantized matrix's input width as `scales.shape(-1) * group_size` and `extract_quantized_matmul_dims` throws `std::invalid_argument` when it disagrees with the activation; `quantized_matmul` crosses the cxx bridge as `UniquePtr<MlxArray>` rather than a `Result`, so that throw is an uncatchable `std::terminate` at the first forward pass, which is exactly the load-time-not-inference-time contract this file exists to hold. The same gap covered the embedding table and the untied head, where `validate_table` skipped its column check entirely once the table was quantized. `validate_quantized_scales` now checks the scales rows, reconstructs the input width the way MLX does and rejects a disagreement, and requires the affine `biases` to have the same shape as the `scales` they are the zero points for (MLX's `validate_quantized_input` throws on that too). `validate_projection` and `validate_table` both route through it. Also cross-checks the dense `.bias` set across `q_proj` / `k_proj` / `v_proj`. `FusedQKVLinear::from_weights_separate` concatenates the three biases only when all three are present and drops the whole set otherwise, the same all-or-nothing rule the affine `biases` check already guarded, with the same silent outcome: the checkpoint loads and the projections that had a bias run without it. Validated by replaying the widened checks over the real `helium-1-preview-2b-4bit` SafeTensors header (559 tensors, all pass), plus five new unit tests: a consistently quantized block that must still be accepted, a mis-packed attention block, a mis-packed output head, zero points that disagree with their scales, and a partial dense bias. The two existing quantized-rejection tests now build scales that describe the model width so they still reach the assertion they were written for. `cargo test --release --lib --features metal,accelerate models::helium`: 27 passed. Clippy and `cargo fmt --all -- --check` clean apart from the two warnings pre-existing on `main`. Refs #837
Security and performance reviewReviewed the surface a prior reviewer did not cover: the fused-path bypass, the quantized loader, the MLX-precondition class, and what the bypass costs. Two findings fixed in 670deff, the rest reported. FixedHIGH — the quantized path never checked the input axis ( Concrete: a HuggingFace repo whose 24 layers ship
Checking against the declared MEDIUM — the dense Validation: replayed the widened checks over the real Verified cleanThe bypass does not weaken bounds checking, and it cannot be evaluated inconsistently. On the reverse question: MLX preconditions, checked against the pinned Performance: the bypass costs nothing. Both gates are opt-in and off by default ( Reported, not fixedLOW — no stop token when LOW — Note. Not addressed here per scope: #929 (tree-wide |
PR finalization auditRan the full verification pass requested for this PR (test coverage, doc accuracy, Regression tests for both fix commits, verified by revertingHIGH ( Went further to reproduce the abort this check exists to prevent. A shape probe that only adds mismatched MEDIUM ( Review commit's RoPE relative-floor anchor ( Positive control review. Test coherence (27 tests across three commits)Read all 27 and mapped each to the property it owns. No duplicate layers found. In particular, four bias-related tests that could look redundant at a glance are each guarding a distinct branch: presence-consistency of affine DocumentationThe
|
…ama path The shared Llama attention never read `rope_traditional` from `config.json`, while the reference `ModelArgs` in mlx_lm/models/llama.py declares it and passes it straight into `initialize_rope`. A Llama, Qwen2 or Qwen2.5 checkpoint declaring the key was therefore rotated with the split-half convention where upstream rotates interleaved. That failure is silent: both conventions consume and produce the same shapes, the KV cache stays consistent, and the model emits fluent text out of a mis-rotated attention, so no shape test and no reading of the output can catch it. #930 threaded the flag from `ModelArgs` to every RoPE route but left it `#[serde(skip)]`, writable only by `helium::ModelArgs::to_llama3_args`. This removes the skip. The field now deserializes with a default of `false`, so no existing checkpoint changes: none of the 182 local configs declares the key, and the default reproduces today's graph exactly. An explicit `null` maps to `false` rather than erroring, because `#[serde(skip)]` ignored whatever the key held and deserializing it must not turn a previously-ignored key into a load failure. Blast radius, verified rather than assumed. Nine sites deserialize `llama3::ModelArgs` from JSON and now pick the key up: the generic Llama loader, six VLM `text_config` loaders (Pixtral, LLaVA, SmolVLM/Idefics3, Idefics2, InternVL, and FastVLM which parses the whole config because it keeps text fields at the root), both pipeline stage executors, and the tensor-parallel runtime. None injects the key; the VLM loaders copy only `quantization` into `text_config`, so a wrapper-level key cannot leak into a nested text backbone. `sanitize_config_json`, which four of those sites run first, only rewrites `Infinity` and `NaN` literals and leaves the key intact. Tests pin both config shapes and the sanitizer. `MllamaTextConfig::to_llama3_args` was a real gap, not a scoping decision: it has always deserialized `rope_traditional` but rebuilt a synthetic config from a fixed key list that dropped it, so an `mllama` checkpoint declaring the key contradicted its own config in the self-attention layers. The key is now forwarded. Cross-attention layers are unaffected since they apply no RoPE. The two fused quantized launchers keep being bypassed rather than the cxx bridge being extended, and the rationale recorded by #930 is replaced rather than duplicated. Reading both C++ bodies, neither is a fused kernel: each builds the same MLX graph the Rust fallback builds, op for op, saving roughly eleven cxx crossings per layer per forward. The bypass therefore costs FFI overhead, not throughput, and both launchers are opt-in and off by default. Extending the signature would touch a surface shared with two neighboring launchers that hardcode the same constant, add a parameter that is `false` at every call site in the tree, and could not be validated, since no checkpoint pairs the key with the quantized Llama path. What the bypass lacked was observability, so `Attention::from_weights` now prints a one-time stderr notice when a traditional-RoPE checkpoint is loaded while either variable asked for a fused path. `eprintln!` because `tracing` has no subscriber in the CLI binary. Tensor parallelism stays refused for Helium, with a corrected reason. The flag now reaches every rank (`local_llama_args` clones the parsed args, pinned by a new test that asserts the flag on every rank's every layer), but Helium's convention is fixed in upstream code and its published `config.json` carries no such key, so the TP runtime's direct parse still yields `false`. Lifting the refusal means routing the rank-local config through `to_llama3_args`, which deserializing a key cannot deliver. Testing: new `src/models/llama3_tests.rs` covers the parse contract in both directions plus null and non-boolean, both VLM config shapes, the sanitizer, the flag reaching `Attention`, and the fused-prefill bypass on a genuinely quantized block. Correctness is asserted at the logits, not the shapes: a model built from a JSON config with the key is bit-identical to one with the flag set programmatically and separated from the split-half model by both an absolute and a relative floor, and the batched route is pinned against the single-sequence route. The #930 test that asserted the opposite behavior is replaced by one showing Helium still needs its conversion. `cargo test --release --lib` green for `models::llama3` (11), `models::helium` (27), `models::mllama` (my test; the pre-existing `ragged_real_tile_rows_match_reference_masked_full_rows` failure is #939), `loading::` (220), `distributed::` (1266). clippy and fmt clean apart from two warnings already on main. Closes #931
…ama path (#951) ## Summary The shared Llama attention path never read `rope_traditional` from `config.json`, while the reference `ModelArgs` in [`mlx_lm/models/llama.py`](https://github.com/ml-explore/mlx-lm/blob/main/mlx_lm/models/llama.py) declares `rope_traditional: bool = False` and passes it straight into `initialize_rope`. A Llama, Qwen2 or Qwen2.5 checkpoint declaring the key was rotated split-half where upstream rotates interleaved. The failure is silent: both conventions consume and produce `[batch, heads, seq, head_dim]`, the KV cache stays consistent, and the model emits fluent text out of a mis-rotated attention. No shape test and no reading of the output can catch it. #930 threaded the flag from `ModelArgs` through every RoPE route but deliberately left it `#[serde(skip)]`. This removes the skip. ## What changed - **`src/models/llama3.rs`** — `ModelArgs::rope_traditional` deserializes with a default of `false`. An explicit `null` maps to `false` rather than erroring, via a small `deserialize_with`: the field was `#[serde(skip)]`, which ignored whatever the key held, so a config carrying `"rope_traditional": null` loaded fine and must keep loading. A non-boolean is still rejected. Added `FUSED_CAUSAL_PREFILL_ENV` / `FUSED_QKV_SPLIT_ROPE_ENV` / `FUSED_ROPE_ENV_VARS` so the gate, the notice and the tests read one list, and `report_fused_rope_bypass_once()`, called from `Attention::from_weights`. - **`src/models/llama3_tests.rs`** (new) — the parse contract, the two VLM config shapes, the sanitizer, the flag reaching `Attention`, the logits-level oracle, the batched route, and the fused-prefill bypass on a genuinely quantized block. - **`src/models/mllama/config.rs`** — `to_llama3_args` now forwards `rope_traditional`, plus a test. - **`src/models/helium_tests.rs`** — replaced the test that asserted the old behavior; renamed `attention_carries_the_traditional_flag_only_for_helium`. - **`src/distributed/tensor_parallel/inference.rs`** — rewrote the Helium TP refusal rationale. - **`src/distributed/tensor_parallel/llama_runtime_tests.rs`** — new `tensor_parallel_llama_propagates_rope_traditional_to_every_rank`. - **`docs/environment-variables.md`**, **`docs/supported-models.md`** — the two fused-path rows now state the bypass; the Helium entry's TP claim was wrong after this change and is corrected. ## The nine deserialization sites Each was read, not assumed. All pick the key up from wherever the checkpoint declares it and default to `false` when it is absent. | Site | Shape it parses | Behavior with the default | |---|---|---| | `src/loading/mod.rs:304` | full config, string-sanitized first | unchanged | | `src/loading/vlm_pixtral.rs:236` | `text_config` sub-object | unchanged | | `src/loading/vlm_llava.rs:173` | `text_config` sub-object | unchanged | | `src/loading/vlm_smolvlm.rs:189` | `text_config` sub-object | unchanged | | `src/loading/vlm_idefics2.rs:170` | `text_config` sub-object | unchanged | | `src/loading/vlm_internvl.rs:125` | `text_config` sub-object | unchanged | | `src/loading/vlm_fastvlm.rs:169` | full config (FastVLM keeps text fields at the root) | unchanged | | `src/distributed/pipeline/stage_executor/{llama,mistral}.rs` | full config, string-sanitized first | unchanged | | `src/distributed/tensor_parallel/llama_runtime.rs:70` | full config, string-sanitized first | unchanged | Two facts that mattered and are now pinned by tests. First, no site injects the key: the six VLM loaders copy only `quantization` into `text_config`, so a wrapper-level `rope_traditional` cannot leak into a nested text backbone, which is correct because in a VLM config the top level describes the multimodal wrapper and not the decoder. Second, `sanitize_config_json`, which four of these sites run before parsing, only rewrites `Infinity` and `NaN` literals by string substitution and leaves the key intact. ## `MllamaTextConfig::to_llama3_args`: a gap, not a scoping decision It builds a synthetic config from a fixed key list, so a key it does not name cannot reach the shared decoder. `MllamaTextConfig` has always deserialized `rope_traditional` with `#[serde(default)]`, and the conversion dropped it, so an `mllama` checkpoint declaring the key rotated split-half in its self-attention layers against what its own config said. Fixed by naming the key. Cross-attention layers are unaffected: they attend to vision features and apply no RoPE. ## Decision on the fused launchers: keep bypassing `fused_qkv_project_split_rope` and `fused_causal_prefill_attention` still hardcode `traditional = false` in C++, and `Attention::forward` still routes a traditional-RoPE model around both. The #930 rationale comment is replaced, not supplemented. Reasons, in order of weight: 1. **Neither is a fused kernel.** Reading both C++ bodies, each is `quantized_matmul`, three `slice`s, `reshape`, `transpose`, `fast::rope`, which is the same MLX graph the Rust fallback builds op for op. What they save is roughly eleven cxx crossings per layer per forward. The "silent performance cliff" argument for extending the bridge therefore describes FFI call overhead, not throughput, and it is invisible next to the quantized matmuls on either side of it. 2. **Both are opt-in and off by default**, so nothing shipped today loses anything. 3. **The signature is shared.** `fused_qkv_project_and_rope` and `fused_qkv_project_split_norm_rope` hardcode the same constant and serve other families. Changing two of four leaves an inconsistent surface; changing four adds a parameter that is `false` at every call site in the tree and pulls families outside this fix into the blast radius. 4. **It could not be validated where it matters.** No checkpoint pairs `rope_traditional` with the quantized Llama path, so an extended bridge would ship on synthetic evidence only. What the bypass actually lacked was observability, so that is fixed rather than argued away: `Attention::from_weights` prints a one-time stderr notice when a traditional-RoPE checkpoint is loaded while either variable asked for a fused path, naming the variables and saying the graph path is used instead. `eprintln!` because `tracing` has no subscriber in the `mlxcel` CLI binary. Revisit if either launcher is promoted to default-on or grows a real kernel; the notice makes both visible. `a_traditional_rope_block_is_routed_around_the_fused_prefill_launcher` asserts the chosen behavior on a genuinely quantized block, in both halves: the flag makes the launcher request a no-op, and the same request on the same weights with the flag off produces a visibly different rotation, so the equality is not vacuous. `the_fused_qkv_rope_launcher_cannot_express_traditional_rope` (from #930) continues to pin the other launcher. ## Tensor parallelism: refusal kept, rationale corrected The old reason ("`rope_traditional` is deliberately not deserializable") is no longer true. The flag now reaches every rank: the TP runtime parses `config.json` into the shared args and `local_llama_args` clones them, and `tensor_parallel_llama_propagates_rope_traditional_to_every_rank` asserts the flag on every rank's every layer rather than inferring it from the `args.clone()`. The refusal survives for a different reason. Helium's convention is fixed in upstream code, not in its config: upstream builds `nn.RoPE(..., traditional=True)` and the published `config.json` carries no `rope_traditional` key at all, which `helium_still_needs_the_conversion_because_its_config_omits_the_key` pins. The TP runtime parses `config.json` directly and never goes through `to_llama3_args`, so a sharded Helium would still parse `false`. Lifting the refusal means routing the rank-local config through the conversion, which is a change to how the runtime builds its config and not something deserializing a key can deliver. ## The #930 test that had to go `a_llama_config_cannot_turn_on_traditional_rope_through_json` parsed a Llama config carrying the key and asserted the parsed value was `false`, which is the exact inverse of this issue's first acceptance criterion. It was replaced (not made to pass) by `helium_still_needs_the_conversion_because_its_config_omits_the_key`, which asserts the thing that is actually still true and is the reason `to_llama3_args` is not now redundant. `attention_carries_the_traditional_flag_only_for_helium` was renamed to `..._from_its_args`, since "only for Helium" stopped being accurate. ## On the token-level acceptance criterion The issue asks for a token-level comparison against a reference on a checkpoint that sets `rope_traditional`. No such public Llama, Qwen2 or Qwen2.5 checkpoint exists, and none of the 182 local configs declares the key, so there is nothing to generate from. The strongest available substitute is implemented instead, and it is a logits comparison rather than a shape check: a model built from a JSON config carrying the key is bit-identical (`max_abs_diff == 0.0`) to one with the flag set programmatically, and separated from the split-half model by both an absolute and a relative floor. That proves the key selects exactly the interleaved rotation and nothing else. What a real-checkpoint run would add is that mlxcel's interleaved rotation matches the reference's, and that equivalence is `fast_rope(traditional = true)` itself, which #930 validated token-exactly against mlx-lm on a real Helium checkpoint. Byte-identical output for existing Llama / Qwen2 / Qwen2.5 checkpoints follows from the default: the parsed args are identical, so the graph is identical. Real-model regression across the local checkpoints is left to the orchestrator. ## Test plan - [x] `cargo check --release --lib --tests --features metal,accelerate` clean (two warnings already on `main`) - [x] `cargo test --release --lib --features metal,accelerate models::llama3` — 11 passed - [x] `cargo test --release --lib --features metal,accelerate models::helium` — 27 passed - [x] `cargo test --release --lib --features metal,accelerate models::mllama` — new test passes; `ragged_real_tile_rows_match_reference_masked_full_rows` fails on `main` too (#939) - [x] `cargo test --release --lib --features metal,accelerate loading::` — 220 passed - [x] `cargo test --release --lib --features metal,accelerate distributed::` — 1266 passed - [x] `cargo test --release --lib --features metal,accelerate server::startup` — 54 passed - [x] `cargo clippy --release --lib --tests --features metal,accelerate` clean - [x] `cargo fmt --all -- --check` clean - [x] `python3 scripts/ci/check_cross_repo_refs.py` clean Closes #931
Summary
Adds Kyutai Helium (
helium) as a fully integrated text model: config parse and validation, loader, detection, registration,LoadedModeldispatch, docs and unit tests. Helium is a dense Llama-shaped decoder (RMSNorm before attention and before the MLP, grouped-query attention, a SwiGLU MLP overgate_proj/up_proj/down_proj, no QK-norm, no MoE, no sliding window, and Llama's own weight key names), so the decoder block, attention and MLP come fromsrc/models/llama3.rsunchanged rather than being copied. Upstream reference: https://github.com/ml-explore/mlx-lm/blob/main/mlx_lm/models/helium.pyThe one architectural difference, and how the fused fast paths were resolved
Upstream builds
nn.RoPE(head_dim, traditional=True, base=rope_theta). Helium rotates interleaved channel pairs(2i, 2i+1); every other Llama-family model in this tree rotates the split-half pairs(i, i + dims/2). The two produce identically shaped tensors from identical weights, so running the wrong one is a silent quality regression that no shape assertion, no cache assertion and no logits-shape assertion can catch.Three code paths in
llama3::Attentioncan rotate Q and K, and only one of them can express the convention:fused_causal_prefill_attention(quantized prefill, opt-in throughMLXCEL_ENABLE_FUSED_CAUSAL_PREFILL_ATTENTION),FusedQKVLinear::forward_split_rope(quantized projection, opt-in throughMLXCEL_ENABLE_FUSED_QKV_SPLIT_ROPE),fast_ropedirectly.forward_split_ropecannot express traditional RoPE. It forwards toffi::fused_qkv_project_split_rope, whose C++ body insrc/lib/mlxcel-core/cpp/mlx_cxx_bridge.cppcallsmlx::core::fast::rope(q, rope_dims, false, rope_base, 1.0f, cache_offset)withtraditionalhardcoded tofalse, and neither the Rust helper nor the bridge signature takes a flag.fused_causal_prefill_attentiondoes the same thing at line 4079 of the same file. Both are reachable only for quantized weights, which is exactly what the Helium validation checkpoint is.Resolution chosen: bypass, not extend. Both fused branches are now gated on
!self.rope_traditional, so a traditional-RoPE model always takes the graph fallback, which receives the flag. Extending the launchers would be the better long-term fix, but it is a cxx bridge signature change that would touch the FFI surface every Llama-family and Gemma-family model shares, and correctness for this port must not wait on it. Neither fused path is enabled by default (both require an environment variable), so the bypass costs the existing families nothing that is on today, and costs Helium nothing at all.How the flag is proved to be honored on every route
Every route Helium can reach was enumerated and threaded, not just the obvious one:
fast_ropeinAttention::forwardself.rope_traditionalfast_rope_batchedinAttention::forward_split_attentionself.rope_traditionalforward_split_attention, RoPE applied before the paged dispatchFusedQKVLinear::forward_split_ropefused_causal_prefill_attentionllama_runtimeper-rankllama3::ModelArgsVerified by mutation, not just by passing.
fast_ropecalls to a literalfalsefailed exactlyhelium_logits_differ_from_the_same_weights_with_split_half_ropeandevery_rope_route_honors_the_flag_consistently, with the other 19 tests green.fast_rope_batchedcalls failed exactlyevery_rope_route_honors_the_flag_consistently, with 20 green. That is the test whose whole purpose is to catch a flag honored on one route and dropped on another, and it does.the_fused_qkv_rope_launcher_cannot_express_traditional_ropebuilds real 4-bit quantized weights, runsforward_split_rope_quantized, and asserts its Q matches the split-half rotation to within 1e-4 while differing from the traditional rotation by more than 1e-3. The fact the bypass rests on is asserted rather than assumed, so if someone later teaches the launcher the flag, that test fails and points at the gate that should then be removed.Existing families stay behaviorally identical
llama3::ModelArgs::rope_traditionalis#[serde(skip)]on purpose. A checkpoint author must not be able to change how an existing Llama checkpoint decodes by adding a config key, so the flag is only ever set programmatically, by the Helium loader.a_llama_config_cannot_turn_on_traditional_rope_through_jsonpins that: a Llamaconfig.jsoncontaining"rope_traditional": truestill parses tofalse.Noted while reading upstream, deliberately out of scope: mlx-lm's own
llama.pydoes haverope_traditional: bool = Falsein itsModelArgsand passes it tonn.RoPE, and mlxcel's Llama path has never read it. Wiring it would be a behavior change for any Llama checkpoint that ships the key, which is not this issue.Every pre-existing family keeps
rope_traditional = false.tensor_parallel_llama_matches_full_model_logitsandtensor_parallel_qwen2_matches_full_model_logits, which compare the sharded runtime against the fullLlama3Model, both still pass.Tensor parallelism is refused on purpose
The TP runtime builds its per-rank model by parsing
config.jsonstraight intollama3::ModelArgs(llama_runtime.rs,local_llama_args), andrope_traditionalis not deserializable, so a sharded Helium would silently rotate split-half while the single-process path rotates interleaved. TheModelType::Helium => "helium"arm ininference.rsexists only to keep the dispatch table total;runtime_kind_forhas no arm forModelType::Heliumand so returnsNone, whichvalidate_supported_runtimeturns into an unsupported-architecture error before any TP load is attempted. Enabling TP here means threading the flag through the rank-local config first, and the comment at the arm says so.Untrusted-config hardening, built in from the start
ModelArgs::validaterejects, at load:num_attention_heads == 0, before any divisibility check, because0.is_multiple_of(0)is true andhead_dim()would then divide by zero.hidden_size,num_hidden_layers,intermediate_size,vocab_size,max_position_embeddings, beforenum_hidden_layerssizes theVec::with_capacityinLlama3Model::from_weightsand before any derivedas i32cast can truncate negative.num_attention_headsis not divisible bynum_key_value_heads(this checkpoint has 20 and 20, so MHA, but the check is not conditional on that).head_dimthat disagrees withhidden_size / num_attention_heads. Authoritative source is the derived value, because upstreamHeliumAttentioncomputesargs.hidden_size // n_headsand never reads the field even though its dataclass requires it. A disagreeing config is rejected rather than one being silently preferred, since a checkpoint built for the other value would mis-shape the attention reshape.mlx::core::fast::roperequiresdimsto be positive, even and no larger than the last axis, and enforces that by throwingstd::invalid_argument;fast_ropecrosses the cxx bridge asUniquePtrrather thanResult, so that throw is an uncatchablestd::terminateat the first forward pass, long after the checkpoint appeared to load. Helium rotates the full head, sodims == head_dimand the third condition holds by construction, but the first two are config-controlled throughhidden_sizeandnum_attention_heads.rope_theta, which RoPE exponentiates per channel and which would otherwise NaN every rotated channel with nothing throwing at all.rms_norm_eps.fast::rms_normnever inspectseps, so a bad value produces NaN hidden states with no error: the checkpoint loads, generation runs, and the output is uniform garbage. This checkpoint's1e-08is unusually small for the family and is explicitly covered as a value that must stay accepted.quantizationblock whosebitsfalls outside1..=32or whosegroup_sizeis non-positive. This is a range check, not an allowlist, because mlxcel deliberately re-derives an effective bit width from the tensor shapes when the declared one disagrees, and an allowlist would reject the mixed-precision exports that behavior serves. This is the 4-bit checkpoint in the set, so it is the one that actually exercises quantized paths.validate_weightsruns beforeLlama3Model::from_weights, not after, becauseFusedQKVLinear::from_weights_separateconcatenatesq_proj/k_proj/v_proj(and theirscalesandbiases) along axis 0 with no shape check of its own, andAttention::forwardthen reshapes the result with config-derived head counts. Both a mismatchedconcatenateand a mismatchedreshapethrow inside MLX, which is an uncatchable abort rather than a load error. It rejects:config.json'svocab_sizeclaims. MLX's gather adds the axis size to a negative index but performs no range check on a positive one, so an id past the last row reads whatever follows the table in the buffer and the result reaches the logits. A table with more rows than claimed is accepted, which is how a vocabulary-padded head is stored.q_proj,k_projandv_projdisagree on whether they are quantized, since the fused loader decides fromq_proj.scalesalone and then concatenates all three.q_proj/k_proj/v_projcarry affinebiases. The fused loader keeps the set only when all three have one and silently drops the whole set otherwise, which dequantizes the survivors without their zero points. Nothing throws; this is the quietest failure in the file.What changed
src/models/helium.rs(new):ModelArgswithvalidate/validate_rope/validate_norm_eps/validate_quantization,TokenIdField,to_llama3_args, thevalidate_weightsshape pass,HeliumModelwrappingLlama3Model, and theLanguageModelimpl.eos_token_idscomes fromconfig.json(2,</s>) rather than delegating, because Helium'stokenizer_config.jsondeclares neither an EOS token nor a chat template andLlama3Model::eos_token_idsreturns Llama 3's hardcoded[128001, 128009], which are outside Helium's 48000-entry vocabulary and would therefore never match, so generation would only ever stop at the token limit.src/models/helium_tests.rs(new): 22 tests. Config surface against the real field set, the head-width authority question, all five validation guards including their accepted boundaries, the real checkpoint's shape signature pinned arithmetically (so it costs nothing where materializing 559 tensors would cost hundreds of megabytes), the four RoPE-convention tests, and six weight-contract rejections.src/models/llama3.rs:ModelArgs::rope_traditional(#[serde(skip)]),Attention::rope_traditional, the flag threaded tofast_ropeandfast_rope_batched, both fused branches gated on it, and the routing argument written out in theAttention::forwarddoc comment.forward_split_attention's// Used by:comment updated perdocs/code-guidelines.md.src/distributed/tensor_parallel/llama_runtime_tests.rs: the two test-onlyLlamaModelArgsstruct literals gainrope_traditional: false.ModelArgsis built through serde or clone-and-mutate everywhere else in the tree, so these were the only construction sites affected. The struct-literal form into_llama3_argsis deliberate for the same reason in reverse: a future field added to the shared config fails Helium to compile rather than silently defaulting.src/models/mod.rs: module declaration,HeliumModelre-export, andModelType::Heliumin all four places that must agree (the enum,ALL_MODEL_TYPES, themetadata()display table, and theall_variants!exhaustiveness list).src/models/detection.rs:"helium" => Ok(ModelType::Helium).src/models/detection_tests.rs: a detection test asserting it does not fall through to the Llama arm, which matters because Helium's config is field-for-field a Llama config apart frommodel_typeand the smallerrms_norm_eps.src/model_metadata.rs: standardconfig_backedregistration.src/loaded_model.rs: variant anddelegate_language_model!arm.src/distributed/tensor_parallel/inference.rs: arch string with the refusal rationale.docs/supported-models.md: the Helium entry under "Text and hybrid model families".Test plan
cargo fmt --all -- --checkcleancargo clippy --release --lib --tests --features metal,accelerate: only the two warnings pre-existing onmaininsrc/multimodal/host_preprocessor.rsandhost_preprocessor_export.rscargo test --release --lib --features metal,accelerate models::helium: 22 passed, 0 failedcargo test --release --lib --features metal,accelerate detection: 78 passed, 0 failedcargo test --release --lib --features metal,accelerate distributed::tensor_parallel::llama_runtime::tests: 19 passed, 0 failed, 4 ignored, including the llama and qwen2 full-model parity testscargo test --release --lib --features metal,accelerate model_metadata: 8 passed, includingevery_variant_is_registered_for_archmodels/helium-1-preview-2b-4bit, promptThe capital of France is, 40 tokens,--no-chat-template(run by the maintainer after review)Closes #837