Skip to content

mlx: support use_bias, qk_norm=none, and frozen running stats - #523

Open
matdou wants to merge 1 commit into
google-research:masterfrom
matdou:mlx-torch-parity-fix-2
Open

matdou wants to merge 1 commit into
google-research:masterfrom
matdou:mlx-torch-parity-fix-2

Conversation

@matdou

@matdou matdou commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

Follow-up to #509. I flagged three more spots where mlx and torch disagree on config handling, but left them out of scope since they just fail.

use_bias=True and qk_norm="none" change which weights exist (extra bias terms, missing RMSNorm params), so a checkpoint using either one crashed on load_safetensors with a key mismatch. use_frozen_running_stats=True was rejected outright, from_hf_config raised NotImplementedError.

All three now work the same way in both backends. use_bias wires into every Linear in ResidualBlock and the attention/FFN layers. qk_norm="none" skips building the query/key RMSNorm submodules entirely, same as torch. Frozen running stats clamps the RevIN mean/std forward from the context boundary instead of letting it keep updating into the horizon.

Checked against real torch weights (~1e-6 match) and against a real TimesFM3Torch.to_dict(), to make sure from_hf_config still parses the production checkpoint shape correctly with and without these fields set.

These three used to just break: use_bias/qk_norm changed which weights
exist so loading crashed, and frozen running stats was flat-out
rejected. Now mlx matches torch on all three, checked against real
weight transplants (~1e-6) and a real checkpoint config shape.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant