Skip to content

feat(models): add Kyutai Helium (helium) text model support - #930

Merged
inureyes merged 3 commits into
mainfrom
feature/issue-837-helium-model
Jul 27, 2026
Merged

inureyes merged 3 commits into
mainfrom
feature/issue-837-helium-model

Conversation

@inureyes

Copy link
Copy Markdown
Member

Summary

Adds Kyutai Helium (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 before attention and before the MLP, grouped-query attention, a SwiGLU MLP over gate_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 from src/models/llama3.rs unchanged rather than being copied. Upstream reference: https://github.com/ml-explore/mlx-lm/blob/main/mlx_lm/models/helium.py

The 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::Attention can rotate Q and K, and only one of them can express the convention:

  1. fused_causal_prefill_attention (quantized prefill, opt-in through MLXCEL_ENABLE_FUSED_CAUSAL_PREFILL_ATTENTION),
  2. FusedQKVLinear::forward_split_rope (quantized projection, opt-in through MLXCEL_ENABLE_FUSED_QKV_SPLIT_ROPE),
  3. the graph fallback, which calls fast_rope directly.

forward_split_rope cannot express traditional RoPE. It forwards to ffi::fused_qkv_project_split_rope, whose C++ body in src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp calls mlx::core::fast::rope(q, rope_dims, false, rope_base, 1.0f, cache_offset) with traditional hardcoded to false, and neither the Rust helper nor the bridge signature takes a flag. fused_causal_prefill_attention does 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:

Route Entry point Handling
Single-sequence prefill and decode fast_rope in Attention::forward receives self.rope_traditional
Batched decode and batched/padded prefill fast_rope_batched in Attention::forward_split_attention receives self.rope_traditional
Paged decode same forward_split_attention, RoPE applied before the paged dispatch covered by the above
Quantized fused projection FusedQKVLinear::forward_split_rope bypassed, cannot express the flag
Quantized fused causal prefill fused_causal_prefill_attention bypassed, cannot express the flag
Tensor parallel llama_runtime per-rank llama3::ModelArgs TP deliberately refused, see below

Verified by mutation, not just by passing.

  • Reverting the single-sequence fast_rope calls to a literal false failed exactly helium_logits_differ_from_the_same_weights_with_split_half_rope and every_rope_route_honors_the_flag_consistently, with the other 19 tests green.
  • Reverting only the fast_rope_batched calls failed exactly every_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.
  • Both mutations were reverted and the suite re-run green.
  • the_fused_qkv_rope_launcher_cannot_express_traditional_rope builds real 4-bit quantized weights, runs forward_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_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. a_llama_config_cannot_turn_on_traditional_rope_through_json pins that: a Llama config.json containing "rope_traditional": true still parses to false.

Noted while reading upstream, deliberately out of scope: mlx-lm's own llama.py does have rope_traditional: bool = False in its ModelArgs and passes it to nn.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_logits and tensor_parallel_qwen2_matches_full_model_logits, which compare the sharded runtime against the full Llama3Model, both still pass.

Tensor parallelism is refused on purpose

The TP runtime builds its per-rank model by parsing config.json straight into llama3::ModelArgs (llama_runtime.rs, local_llama_args), and rope_traditional is not deserializable, so a sharded Helium would silently rotate split-half while the single-process path rotates interleaved. The ModelType::Helium => "helium" arm in inference.rs exists only to keep the dispatch table total; runtime_kind_for has no arm for ModelType::Helium and so returns None, which validate_supported_runtime turns 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::validate rejects, at load:

  • num_attention_heads == 0, before any divisibility check, because 0.is_multiple_of(0) is true and head_dim() would then divide by zero.
  • Out-of-range hidden_size, num_hidden_layers, intermediate_size, vocab_size, max_position_embeddings, before num_hidden_layers sizes the Vec::with_capacity in Llama3Model::from_weights and before any derived as i32 cast can truncate negative.
  • An indivisible head split, and a GQA split where num_attention_heads is not divisible by num_key_value_heads (this checkpoint has 20 and 20, so MHA, but the check is not conditional on that).
  • A declared head_dim that disagrees with hidden_size / num_attention_heads. Authoritative source is the derived value, because upstream HeliumAttention computes args.hidden_size // n_heads and 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.
  • An odd or non-positive head width. mlx::core::fast::rope requires dims to be positive, even and no larger than the last axis, and enforces that by throwing std::invalid_argument; fast_rope crosses the cxx bridge as UniquePtr rather than Result, so that throw is an uncatchable std::terminate at the first forward pass, long after the checkpoint appeared to load. Helium rotates the full head, so dims == head_dim and the third condition holds by construction, but the first two are config-controlled through hidden_size and num_attention_heads.
  • A non-positive or non-finite rope_theta, which RoPE exponentiates per channel and which would otherwise NaN every rotated channel with nothing throwing at all.
  • A non-finite, negative or zero rms_norm_eps. fast::rms_norm never inspects eps, so a bad value produces NaN hidden states with no error: the checkpoint loads, generation runs, and the output is uniform garbage. This checkpoint's 1e-08 is unusually small for the family and is explicitly covered as a value that must stay accepted.
  • A quantization block whose bits falls outside 1..=32 or whose group_size is 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_weights runs before Llama3Model::from_weights, not after, because FusedQKVLinear::from_weights_separate concatenates q_proj / k_proj / v_proj (and their scales and biases) along axis 0 with no shape check of its own, and Attention::forward then reshapes the result with config-derived head counts. Both a mismatched concatenate and a mismatched reshape throw inside MLX, which is an uncatchable abort rather than a load error. It rejects:

  • An embedding table or output head with fewer rows than config.json's vocab_size claims. 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.
  • Any projection, bias or norm whose shape disagrees with the config, in every layer rather than layer 0. The quantized path does not skip the row check: packing compresses the input axis only, so the output row count is still the width the fused concatenation and the attention reshape depend on. Skipping it on the quantized path is the same carve-out that let K and V be sliced from arbitrary interior channels in the previous PR in this chain.
  • An attention block whose q_proj, k_proj and v_proj disagree on whether they are quantized, since the fused loader decides from q_proj.scales alone and then concatenates all three.
  • An attention block where only some of q_proj / k_proj / v_proj carry affine biases. 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): ModelArgs with validate / validate_rope / validate_norm_eps / validate_quantization, TokenIdField, to_llama3_args, the validate_weights shape pass, HeliumModel wrapping Llama3Model, and the LanguageModel impl. eos_token_ids comes from config.json (2, </s>) rather than delegating, because Helium's tokenizer_config.json declares neither an EOS token nor a chat template and Llama3Model::eos_token_ids returns 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 to fast_rope and fast_rope_batched, both fused branches gated on it, and the routing argument written out in the Attention::forward doc comment. forward_split_attention's // Used by: comment updated per docs/code-guidelines.md.
  • src/distributed/tensor_parallel/llama_runtime_tests.rs: the two test-only LlamaModelArgs struct literals gain rope_traditional: false. ModelArgs is built through serde or clone-and-mutate everywhere else in the tree, so these were the only construction sites affected. The struct-literal form in to_llama3_args is 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, HeliumModel re-export, and ModelType::Helium in all four places that must agree (the enum, ALL_MODEL_TYPES, the metadata() display table, and the all_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 from model_type and the smaller rms_norm_eps.
  • src/model_metadata.rs: standard config_backed registration. src/loaded_model.rs: variant and delegate_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 -- --check clean
  • cargo clippy --release --lib --tests --features metal,accelerate: only the two warnings pre-existing on main in src/multimodal/host_preprocessor.rs and host_preprocessor_export.rs
  • cargo test --release --lib --features metal,accelerate models::helium: 22 passed, 0 failed
  • cargo test --release --lib --features metal,accelerate detection: 78 passed, 0 failed
  • cargo 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 tests
  • cargo test --release --lib --features metal,accelerate model_metadata: 8 passed, including every_variant_is_registered_for_arch
  • Mutation testing of both RoPE routes, described above
  • Real-checkpoint token-exact validation against models/helium-1-preview-2b-4bit, prompt The capital of France is, 40 tokens, --no-chat-template (run by the maintainer after review)

Closes #837

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
@inureyes inureyes added status:review Under review type:enhancement New features, capabilities, or significant additions priority:low Low priority area:models Model architectures, weights, loading, metadata labels Jul 26, 2026
@inureyes

Copy link
Copy Markdown
Member Author

Real-checkpoint validation on Apple M1 Ultra (Metal + Accelerate, release build).

Token-exact against the mlx-lm reference for 40 greedy tokens on models/helium-1-preview-2b-4bit: mlxcel generate -m models/helium-1-preview-2b-4bit -p "The capital of France is" -n 40 --no-chat-template reproduces the reference continuation exactly, starting Paris. It is the largest city in the country, at 190.6 tok/s. mlxcel arch reports Kyutai Helium (dense Llama shape, traditional RoPE).

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: fused_qkv_project_split_rope hardcodes traditional = false in its C++ body, so a port that routed through it would have applied split-half rotation to a model expecting the traditional convention. That failure mode produces fluent, plausible text and correctly-shaped tensors, so it is invisible to shape tests and to eyeballing; reproducing the reference token for token on the quantized checkpoint is what rules it out.

Worth noting for the reviewers: the PR bypasses the fused launchers when rope_traditional is set rather than extending the cxx bridge signature, and both launchers are env-var opt-in and off by default, so the bypass costs nothing on the default path today.

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
@inureyes

Copy link
Copy Markdown
Member Author

Implementation Review Summary

Intent

Add Kyutai Helium (helium) as a fully integrated text model, reusing the shared dense Llama decoder and threading the one architectural difference (traditional/interleaved RoPE) through every route the family can reach.

Findings Addressed

  • Attention::forward and Attention::from_weights were modified without the // Used by: comment docs/code-guidelines.md requires (HIGH). Both are shared functions this PR changed the behavior of: the two !self.rope_traditional gates live in the former, the flag wiring in the latter. All three of the guideline's update triggers fired (attention implementation, behavior change, new model reusing it). src/models/llama3.rs:149, src/models/llama3.rs:539.
  • The // Used by: list the PR did update, on forward_split_attention, named only Llama3, Qwen2 and Helium (MEDIUM). The real set that reaches this attention also includes the mllama text decoder's self-attention layers, seven VLM text backbones built on Llama3Model / Qwen2Model (Pixtral, LLaVA, SmolVLM/Idefics3, Idefics2, InternVL, FastVLM, dots.ocr), the llama and mistral pipeline stage executors, and the tensor-parallel Llama runtime. src/models/llama3.rs:348.
  • helium_logits_differ_from_the_same_weights_with_split_half_rope compared against a bare absolute 1e-3 with no anchor to the logits magnitude (MEDIUM). Independently measured: gap 1.19e-2 against a logits scale of 6.89, so ~12x margin today, but a future change that widens tiny_args or rescales filled could grow the logits without growing the separation and drift the test back toward the blindness an earlier filler already caused once. Added a second, relative floor (gap > logits_scale * 1e-4, ~17x margin), both measured values in the failure messages, and a comment pinning why both floors exist. src/models/helium_tests.rs:367.

Remaining Items

  • helium::dim_eq duplicates gpt2::dim_eq byte for byte (LOW) — deliberately left. The non-reuse of gpt2::validate_embedding_table is sound (it takes an already-built UnifiedEmbedding, which only exists after Llama3Model::from_weights has run ffi::concatenate, i.e. past the point Helium must reject). dim_eq is a weaker case, but the correct fix is hoisting it to a shared module, which would touch gpt2 / gpt_bigcode / gpt_neox and is out of scope for a model-port PR. src/models/helium.rs:453.
  • max_position_embeddings is validated but never consumed (LOW) — llama3::ModelArgs has no such field, so to_llama3_args drops it. Harmless, and cheap insurance if it is ever wired. src/models/helium.rs:262.

Verified independently

  • Shared-file inertness. rope_traditional is #[serde(skip)], so every JSON-parsed config yields false. Only three construction sites of llama3::ModelArgs exist tree-wide: the two test-only struct literals the PR updated, and mllama/config.rs:77, which builds through serde_json::from_value and therefore also gets false. Both fused gates are inert at false (identical predicate and identical forward_split_rope call as main), and the let fused_split_rope = ... refactor is semantically identical. tensor_parallel_llama_matches_full_model_logits and tensor_parallel_qwen2_matches_full_model_logits pass.
  • Fused-launcher test really exercises the FFI. forward_split_rope_quantized (layers.rs:1772) is the ungated twin of forward_split_rope and calls the same ffi::fused_qkv_project_split_rope; there is no Rust fallback. The C++ hardcode is confirmed at mlx_cxx_bridge.cpp:3872-3873.
  • TP refusal is real and doubly guarded. runtime_kind_for has no Helium arm and returns None, so validate_supported_runtime bails; even past that, load_model_with_tensor_parallel's match has no Helium arm and bails. The fallback_architecture arm only keeps that match total.
  • head_dim authority. Upstream's HeliumAttention computes hidden_size // n_heads and never reads the field; the port rejects a disagreeing config rather than preferring one, and passes the derived value onward.
  • Config-trust class. RoPE dims evenness and positivity are validated before fast_rope (and dims == head_dim == the last axis by construction); gather bounds come from real tensor shapes via validate_table, not config fields; zero checks precede divisibility checks; rms_norm_eps guard accepts this checkpoint's legitimate 1e-08; GQA divisibility checked. num_attention_heads and num_key_value_heads are transitively bounded by the hidden_size ceiling, so no as i32 cast can truncate. The PR adds a quantization range check rather than new exposure under fix(core): an unvalidated config.json quantization block reaches MLX and terminates the process #929.

Not this PR

models::mllama::text::tests::ragged_real_tile_rows_match_reference_masked_full_rows fails deterministically (3.7252903e-9 against an exact assert_eq!(.., 0.0)). Confirmed pre-existing by running it on main in a separate worktree, where it fails with the identical value. It is mllama cross-attention code that never touches llama3::Attention and applies no RoPE. ci.yml runs no cargo test, which is why it went unnoticed. Worth a separate issue.

Verification

  • All stated requirements implemented
  • No placeholder/mock code remaining
  • Integrated into project code flow (mod.rs enum / ALL_MODEL_TYPES / metadata() / all_variants!, detection.rs, model_metadata.rs, loaded_model.rs, inference.rs, docs/supported-models.md)
  • Project conventions followed
  • Existing modules reused where applicable
  • No unintended structural changes
  • Tests pass (models::helium 22, models::detection 25, model_metadata 8, distributed::tensor_parallel 338, models:: 671/672 with the one pre-existing failure above; cargo fmt clean; clippy clean apart from the two warnings pre-existing on main)

…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
@inureyes

Copy link
Copy Markdown
Member Author

Security and performance review

Reviewed 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.

Fixed

HIGH — the quantized path never checked the input axis (src/models/helium.rs, validate_projection / validate_table). The row check ran on the quantized path (good, that was the #926 carve-out), but the input width did not, 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.

Concrete: a HuggingFace repo whose 24 layers ship q_proj/k_proj/v_proj as weight [2560, 160] + scales [2560, 20] (packed for an input width of 1280, not 2560) passes every check in the file. Rows are 2560 ✓, all three agree ✓, and infer_quantization_bits derives 160*32/(20*64) = 4, which equals the declared bits, so the loader does not complain either. At the first token, extract_quantized_matmul_dims (mlx/ops.cpp:139-152 at the pinned b7c3dd6d) computes w_inner_dims = 160*32/4 = 1280, finds it does not match the activation's 2560, and throws std::invalid_argument. quantized_matmul crosses the cxx bridge as UniquePtr<MlxArray>, so that throw is an uncatchable std::terminate: SIGABRT at first inference, which is exactly the load-time-not-inference-time contract this file exists to hold. validate_table had the wider version of the same hole, skipping its column check entirely once the table was quantized, so a mis-packed lm_head aborted in the final quantized_matmul and a mis-packed model.embed_tokens aborted in the first rms_norm (fast.cpp requires the norm weight to equal the last axis).

validate_quantized_scales now reconstructs the input width the way MLX does, scales.shape(-1) * group_size, and rejects a disagreement at load. It also checks the scales rows and requires the affine biases to have the same shape as their scales, which validate_quantized_input (mlx/ops.cpp:97-103) throws on and which presence-only checking missed.

Checking against the declared group_size is deliberate: the affine loader trusts it and re-derives bits from the shapes, and FusedQKVLinear::from_weights_separate does so unconditionally, so the declared value is what reaches MLX. With this check in place w_inner_dims == in_features holds by construction for every layout the loader will accept.

MEDIUM — the dense .bias set had no all-or-nothing check. FusedQKVLinear::from_weights_separate (layers.rs:1583-1596) concatenates the q/k/v .bias tensors only when all three are present and drops the whole set otherwise, the identical rule the PR already guards for the affine .biases, with the identical silent outcome: the checkpoint loads and the projections that had a bias run without one. Now cross-checked alongside has_quant_biases.

Validation: replayed the widened checks over the real helium-1-preview-2b-4bit SafeTensors header (all 559 tensors pass, no false rejection), plus five new unit tests including a positive control so the guard cannot pass by rejecting everything quantized. models::helium 27 passed. Clippy and cargo fmt --all -- --check clean apart from the two host_preprocessor* warnings pre-existing on main.

Verified clean

The bypass does not weaken bounds checking, and it cannot be evaluated inconsistently. rope_traditional is an immutable per-Attention field set once in from_weights, so prefill and decode cannot disagree. Every RoPE application site this family can reach is threaded: llama3.rs has exactly four (fast_rope x2, fast_rope_batched x2); batched decode and paged decode both funnel through forward_split_attention; tensor parallelism is refused by runtime_kind_for; pipeline parallelism is refused by resolve_stage_family's other => bail! arm (stage_executor/mod.rs:374); and the OpenXLA emitter rejects an unknown model_type outright (emitter/config.rs:1002), so there is no second engine that could silently rotate split-half.

On the reverse question: fused_qkv_project_split_rope does carry one shape assertion the Rust fallback lacks (mlx_cxx_bridge.cpp:3856, proj_shape[2] != qkv_cols), and it matters because slice_last_dim bottoms out in MLX slice, which clamps an out-of-range stop rather than throwing (normalize_slice, ops.cpp:726-733). A fused QKV wider than q_out + 2*kv_out would therefore have its trailing channels silently discarded on the fallback path. Helium's own per-projection row checks are strictly stronger than that assertion (they pin each of q, k and v individually, not just their sum), so nothing is lost. fused_causal_prefill_attention has no equivalent assertion at all.

MLX preconditions, checked against the pinned b7c3dd6d. fast::rope requires dims positive, even and <= x.shape(-1) (fast.cpp:407-423) and takes std::log(base) in the fallback / std::log2(base_) in the Metal path (backend/metal/rope.cpp:120), so a zero, negative or non-finite rope_theta NaNs every channel silently: both guarded, and this checkpoint's unusual 100000.0 passes. fast::rms_norm never inspects eps (fast.cpp, the fallback adds it under an rsqrt unexamined) but does require the weight to be 1-D of exactly x.shape(-1): the eps guard is a range check, so 1e-08 stays accepted, and validate_norm pins every norm width. x.ndim() >= 3 holds by construction after the head transpose. No unguarded config-derived argument remains on any MLX entry point this port feeds.

Performance: the bypass costs nothing. Both gates are opt-in and off by default (llama3.rs:199 presence-enables, layers.rs:1732 early-returns None when unset; docs/environment-variables.md:237-238 documents both as off). More to the point, neither "fused" launcher is a fused kernel: mlx_cxx_bridge.cpp:3840-3878 and 4049-4093 build the identical MLX graph the Rust fallback builds (quantized_matmul -> 3 slice -> 3 reshape -> 3 transpose -> 2 rope), just issued from one C++ frame. The entire saving is roughly 11 cxx crossings and UniquePtr allocations per attention layer per forward, which is ~264 per token on this 24-layer model, against a memory-bandwidth-bound 2B decode step. Helium gives up no MLX work at all, and because it skips forward_split_rope before that function's per-call std::env::var lookup, it actually does slightly less per-layer work than a Llama of the same shape. The unfused path materializes nothing extra per decode step; it is the path every Llama-family model runs today.

Reported, not fixed

LOW — no stop token when config.json omits eos_token_id (helium.rs, eos_token_ids). The PR deliberately does not delegate to Llama3Model::eos_token_ids (correct: Llama 3's hardcoded ids are outside Helium's 48000-entry vocabulary). But the fallback is unwrap_or_default(), an empty Vec, so a Helium checkpoint without the key generates to the token cap with no stop condition and no warning. Consider a load-time warning; the published checkpoint has eos_token_id: 2 so nothing is broken today.

LOW — validate_quantization's range is wider than MLX's real precondition. bits in 1..=32 and any positive group_size are accepted, but MLX's affine kernels support bits in {2,3,4,5,6,8} and group_size in {32,64,128}, and infer_quantization_bits short-circuits when the derived width equals the declared one, so a declared bits: 1 or bits: 7 reaches the Metal kernel lookup. This is #929's territory and the PR adds strictly more validation than the tree-wide baseline, so there is no new exposure here; noting it because the guard's doc comment claims to close the abort class it only narrows.

Note. attention_bias and mlp_bias are threaded into llama3::ModelArgs but that struct never reads either field (bias presence is decided from the weight map), and max_position_embeddings is validated but never passed to to_llama3_args. All three are inert, not wrong. Pre-existing, out of scope.

Not addressed here per scope: #929 (tree-wide quantization block) and #931 (Llama ignoring rope_traditional from config.json). No new exposure added to either.

@inureyes

Copy link
Copy Markdown
Member Author

PR finalization audit

Ran the full verification pass requested for this PR (test coverage, doc accuracy, // Used by: staleness, lint/format) and found nothing that needed changing. No commits were added; the branch is unchanged.

Regression tests for both fix commits, verified by reverting

HIGH (670defffa, quantized input-width check in validate_quantized_scales). Neutralized the described != Some(in_features) check in isolation and reran models::helium: exactly loading_rejects_a_quantized_projection_packed_for_a_different_input_width and loading_rejects_a_quantized_output_head_packed_for_a_different_input_width failed (25 passed, 2 failed), both as clean expect_err panics, everything else green. Restored the check and reran clean (27 passed).

Went further to reproduce the abort this check exists to prevent. A shape probe that only adds mismatched .scales/.biases next to an untouched float .weight (the pattern the shipped rejection tests use) does not reach MLX at all; it is intercepted earlier by infer_quantization_bits in layers.rs, an independent self-consistency check between .weight and .scales that returns a clean Result::Err. That check cannot catch the real bug class, though, because a checkpoint honestly packed for a different hidden_size than the one config.json declares is internally self-consistent (weight and scales agree with each other) and only disagrees with the external config value. Built a second probe with genuinely packed weights (mlxcel_core::quantize_weights) for a wrong input width and drove a real forward pass with the guard disabled: libc++abi: terminating due to uncaught exception of type std::invalid_argument: [quantized_matmul] Last dimension of first input with shape (..., 64) does not match the expanded quantized matrix..., process aborted with SIGABRT. That confirms the "uncatchable std::terminate at the first forward pass" claim in the PR description and code comments is not just documentation, it reproduces. Both probes were temporary, run in isolation, and removed before restoring the guard.

MEDIUM (670defffa, dense .bias all-or-nothing check). Neutralized the q.has_dense_bias != k.has_dense_bias || ... check in isolation and reran: exactly loading_rejects_an_attention_block_with_a_partial_dense_bias failed (26 passed, 1 failed), clean panic, everything else green. Restored and reran clean.

Review commit's RoPE relative-floor anchor (8d0b6ec26). This one doesn't change production code, so "revert and observe failure" doesn't apply the same way; instead reproduced the PR's own mutation check: raised the relative floor from logits_scale * 1e-4 to logits_scale * 1e-2 and reran helium_logits_differ_from_the_same_weights_with_split_half_rope alone. It failed with gap 0.0118608475, logits scale 6.8892975, matching the exact values quoted in the commit message. Confirms the assertion is genuine and deterministic, not a no-op.

Positive control review. the_real_checkpoint_shape_signature_matches_what_loading_expects was widened in the security commit with assert_eq!(40 * group_size, args.hidden_size as i32) and the intermediate-size equivalent. This is not tautological: the pre-existing assertions in the same test are truncating-division checks (hidden_size / group_size == 40), which stay green even if group_size didn't divide hidden_size evenly, since integer division floors silently. The new assertions test exact multiplicative reconstruction, groups * group_size == in_features, which is precisely the equality validate_quantized_scales performs against the real checkpoint's own numbers. a_consistently_quantized_attention_block_is_accepted is a separate, code-exercising positive control (it actually calls validate_weights), distinct from this arithmetic pin.

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 .biases across q/k/v, shape-consistency of .biases against .scales for one projection, presence-consistency of dense .bias across q/k/v (the MEDIUM fix), and quantized-vs-float consistency across q/k/v. No overlap.

Documentation

The docs/supported-models.md Helium entry already reflects all three commits, not just the port commit. It explicitly states that a projection or table shape check runs "on both axes and including on the quantized path, where packing compresses the input axis only... and the input width is recovered from the scales the way MLX recovers it" (the HIGH fix) and that an attention block is rejected when q/k/v "disagree on whether they carry a dense bias" (the MEDIUM fix). This was folded into the doc line progressively as each commit landed (+1/-0 net diff across the PR, +1/-1 in the security commit replacing the port commit's line in place), so no further edit was needed. No Korean counterpart of this file exists anywhere in docs/, so no translation question arises.

// Used by: staleness

Across every file this PR touches, only src/models/llama3.rs carries model-family // Used by: comments. The two on Attention::forward and Attention::from_weights are verbatim identical to each other, and forward_split_attention's comment correctly defers to that list instead of restating a shorter one. The two call-site-only "Used by" comments elsewhere in the same file (TransformerBlock::forward_batched, Llama3Model::forward_batched_impl) document their caller function, a different and pre-existing convention, and are unrelated to this PR's changes. No stale or contradictory comment found.

Lint and format

cargo fmt --all -- --check clean. cargo clippy --release --lib --tests --features metal,accelerate shows only the two pre-existing warnings in src/multimodal/host_preprocessor.rs and host_preprocessor_export.rs, unrelated to this PR. python3 scripts/ci/check_cross_repo_refs.py clean.

Final counts

models::helium: 27 passed, 0 failed. models::detection: 78 passed, 0 failed. distributed::tensor_parallel::llama_runtime::tests: 19 passed, 0 failed, 4 ignored (including the Llama and Qwen2 full-model parity tests). model_metadata: 8 passed, 0 failed. The known pre-existing models::mllama::text::tests::ragged_real_tile_rows_match_reference_masked_full_rows failure was not touched and was not re-verified here since it is already tracked separately.

Nothing required a code, test, or doc change, so no new commit was added.

@inureyes inureyes added status:done Completed and removed status:review Under review labels Jul 27, 2026
@inureyes
inureyes merged commit 1f21b1f into main Jul 27, 2026
5 checks passed
@inureyes
inureyes deleted the feature/issue-837-helium-model branch July 27, 2026 01:57
inureyes added a commit that referenced this pull request Jul 28, 2026
…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
inureyes added a commit that referenced this pull request Jul 28, 2026
…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
@inureyes inureyes self-assigned this Aug 31, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:models Model architectures, weights, loading, metadata priority:low Low priority status:done Completed type:enhancement New features, capabilities, or significant additions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

feat(models): add Kyutai Helium (helium) text model support

1 participant