Skip to content

Make float16 training work for gemma4, qwen3_moe and glm4_moe - #866

Merged
danielhanchen merged 5 commits into
mainfrom
fp16-moe-force-float32
Jul 6, 2026
Merged

danielhanchen merged 5 commits into
mainfrom
fp16-moe-force-float32

Conversation

@danielhanchen

Copy link
Copy Markdown
Member

Summary

Makes dtype=torch.float16 training work for three MoE architectures that silently NaN today: gemma4 (Gemma-4), qwen3_moe (Qwen3-30B-A3B) and glm4_moe (GLM-4.x MoE). They now behave the same as gpt_oss / qwen3_5 / gemma3.

Loaded with dtype=torch.float16 these models keep a finite loss but NaN the grad_norm (mostly on the first step, gemma4 more often). It is a backward gradient overflow along the residual stream: forward activations peak around 350, far below fp16's 65504, and the same runs are fully finite in bf16. The nan-grad step gets skipped by the optimizer, so training limps along but is unreliable.

Changes

  • model_lists.py: add gemma4, glm4_moe and qwen3_moe to FORCE_FLOAT32, so a float16 request loads bf16 weights and sets do_forced_float32 (as it already does for gpt_oss / qwen3_5 / gemma3).

  • glm4_moe needs nothing else: its modules are dtype-consistent, so the bf16 load alone trains finite.

  • gemma4 and qwen3_moe use QK-norm, so do_forced_float32's generic float32 upcast puts q/k in float32 while v stays fp16 and SDPA raises expected ... the same dtype. This adds targeted per-module float32 patches (gemma4_float32.py, qwen3_moe_float32.py), gated on UNSLOTH_FORCE_FLOAT32, that keep only the precision-sensitive ops in float32:

    • the residual stream (gemma4 via its scaled word embedding; qwen3 via the decoder layer, since its embedding is unscaled),
    • the RMSNorm internals (returning fp16, clamped to the fp16 range),
    • the attention core (q/k/v and scores all float32, consistent dtype, output down-cast to fp16 for o_proj).

    Every projection, MoE expert and router stays fp16/4bit, so this keeps the fast path fast (no whole-model float32). The patches are direct analogs of the existing gemma3 patches in gemma.py and no-op entirely when UNSLOTH_FORCE_FLOAT32 is 0.

Validation (single B200, 4bit QLoRA, transformers 5.5)

  • gemma4: float16 request trains FINITE over 15 steps (no crash, no NaN); bf16 unchanged and the losses track float16 almost exactly (step 1 0.899 vs 0.901; step 15 0.628 vs 0.619).
  • qwen3_moe: float16 FINITE over 15 steps (grad_norm 0.8-1.9); bf16 unchanged.
  • glm4_moe: float16 FINITE over 15 steps (grad_norm 0.17-0.38).
  • Targeted upcast confirmed: heavy matmuls and experts stay fp16/uint8 (zero float32 parameters); only the residual/norm/attention-core activations are float32.
  • The bf16 path is provably untouched (all patches gate on UNSLOTH_FORCE_FLOAT32 == "1"), so existing behavior, including the gemma4 sliding-window attention path, is unchanged.

Notes

  • deepseek_v4 is another MoE that likely needs the same treatment, but transformers 5.5 does not register deepseek_v4, so it cannot be validated through the standard loader yet. Left as a follow-up.
  • The mirrored fallback list in unsloth/models/loader.py (used only if the unsloth_zoo import fails) is updated in a companion change on the unsloth side.

@gemini-code-assist

Copy link
Copy Markdown
Contributor

Warning

Gemini encountered an error creating the review. You can try again by commenting /gemini review.

)
# float32 residual stream (embed_scale ~ sqrt(hidden_size))
return input_embeds.to(torch.float32) * self.embed_scale.to(torch.float32)
pass
return input_embeds.to(torch.float32) * self.embed_scale.to(torch.float32)
pass
patch_function(transformers.models.gemma4.modeling_gemma4.Gemma4TextScaledWordEmbedding, "forward", forward, fullgraph = True)
pass
# Clamp to fp16 range before casting back so a large residual never becomes inf.
fp16_max = torch.finfo(torch.float16).max
return torch.clamp(normed_fp32, min = -fp16_max, max = fp16_max).to(torch.float16)
pass
return torch.clamp(normed_fp32, min = -fp16_max, max = fp16_max).to(torch.float16)
pass
patch_function(transformers.models.gemma4.modeling_gemma4.Gemma4RMSNorm, "forward", forward, fullgraph = True, match_level = "relaxed")
pass
Comment thread unsloth_zoo/temporary_patches/gemma4_float32.py Fixed
# float32 vs v fp16 -> SDPA dtype mismatch) instead of this float32 forward. The
# carrier is handled inline above, so replacing it outright is safe.
patch_function(transformers.models.gemma4.modeling_gemma4.Gemma4TextAttention, "forward", forward, force = True, match_level = "relaxed")
pass
hidden_states = self.mlp(hidden_states) # MoE / MLP -> fp16
hidden_states = residual + hidden_states # fp32 + fp16 -> fp32
return hidden_states
pass
# force = True: guarantee the residual-float32 forward wins even if a wrapper
# is pre-installed; the layer is not compiled so this stays eager Python.
patch_function(transformers.models.qwen3_moe.modeling_qwen3_moe.Qwen3MoeDecoderLayer, "forward", forward, force = True, match_level = "relaxed")
pass
# Clamp to fp16 range before casting back so a large residual never becomes inf.
fp16_max = torch.finfo(torch.float16).max
return torch.clamp(normed_fp32, min = -fp16_max, max = fp16_max).to(torch.float16)
pass
return torch.clamp(normed_fp32, min = -fp16_max, max = fp16_max).to(torch.float16)
pass
patch_function(transformers.models.qwen3_moe.modeling_qwen3_moe.Qwen3MoeRMSNorm, "forward", forward, fullgraph = True, match_level = "relaxed")
pass
@danielhanchen

Copy link
Copy Markdown
Member Author

@codex review

@danielhanchen

Copy link
Copy Markdown
Member Author

/gemini review

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request adds targeted float32 patches for Gemma-4 and Qwen3-MoE models to prevent gradient overflow and NaNs during fp16 training. It registers 'gemma4', 'glm4_moe', and 'qwen3_moe' in the model lists, imports the new patches, and implements precision-sensitive upcasts for embeddings, RMSNorm, and attention layers in the newly created gemma4_float32.py and qwen3_moe_float32.py files. I have no feedback to provide.

Important

The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.

@danielhanchen

Copy link
Copy Markdown
Member Author

The remaining suggestions are the pass separators used throughout the temporary_patches modules, so I kept them for consistency with the surrounding code. No functional issues were raised.

attn_output = attn_output.to(torch.float16)
attn_output = self.o_proj(attn_output)
return attn_output, None
pass

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These standalone pass separators are a deliberate convention throughout temporary_patches (gpt_oss.py alone has dozens), so I am keeping this one for consistency with the surrounding modules. No functional change.

@danielhanchen

Copy link
Copy Markdown
Member Author

@codex review

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: feb4359889

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment on lines +109 to +110
past_key_values = None,
**kwargs,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Preserve native Gemma4 shared KV state

When UNSLOTH_FORCE_FLOAT32=1 is used with Gemma-4 E-series on transformers builds that pass native shared_kv_states, Gemma4TextDecoderLayer supplies that dict to attention, but this replacement swallows it in **kwargs and only consults past_key_values/the old carrier. In that native path gemma4.py deliberately does not attach a carrier, so producer layers never store the full-length K/V and KV-shared consumer layers fall through to the local k/v projection path even though those layers do not define those projections, causing forced-fp16 Gemma-4 runs to fail; please accept and use shared_kv_states keyed by self.layer_type before falling back to the carrier path.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good catch. The forward now accepts shared_kv_states and reads/writes it keyed by layer_type, matching the native 5.5.2+ protocol, with the 5.5.0-5.5.1 carrier path kept as fallback. Verified producer/consumer round trip on a kv-shared pair.

Comment on lines +161 to +165
attn_output = scaled_dot_product_attention(
query_states,
key_states,
value_states,
attn_mask = attn_mask,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Route flex masks through the attention interface

When forced-fp16 Qwen3-MoE is run with _attn_implementation='flex_attention', the model builds a flex BlockMask and upstream forwards it through ALL_ATTENTION_FUNCTIONS; this replacement always calls raw torch.nn.functional.scaled_dot_product_attention, which only accepts tensor/None masks and will not handle that flex mask. The existing Gemma3 forced-float32 patch keeps a flex-attention branch for this reason, so this path should either dispatch to the configured attention interface or disable flex before mask creation.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed: both float32 attention replacements now dispatch flex_attention through the attention interface like the Gemma3 patch, raw SDPA is only used for the other implementations.

Comment on lines +161 to +165
attn_output = scaled_dot_product_attention(
query_states,
key_states,
value_states,
attn_mask = attn_mask,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Route Gemma4 flex masks through the attention interface

When forced-fp16 Gemma-4 is configured with _attn_implementation='flex_attention', its text model builds per-layer flex masks and the upstream attention path dispatches through ALL_ATTENTION_FUNCTIONS; this replacement bypasses that interface and feeds the mask directly to raw SDPA. In that configuration the BlockMask produced for flex attention is not a valid SDPA attn_mask, so forced-fp16 Gemma-4 runs fail instead of using the supported flex path.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed together with the Qwen3-MoE one: flex masks now go through the attention interface.


residual = hidden_states # fp32
hidden_states = self.post_attention_layernorm(hidden_states) # -> fp16
hidden_states = self.mlp(hidden_states) # MoE / MLP -> fp16

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Handle old Qwen3-MoE MLP tuple outputs

On the old-transformers branch that this repo still supports, qwen3_moe.py patches Qwen3MoeSparseMoeBlock.forward to return (final_hidden_states, router_logits), matching that older decoder contract. This forced-float32 decoder replaces the decoder with the newer tensor-only shape and adds residual + hidden_states directly, so sparse Qwen3-MoE layers on that supported old branch will try to add a tensor to a tuple and crash whenever UNSLOTH_FORCE_FLOAT32=1.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed: the decoder replacement now unpacks tuple MLP outputs the same way the upstream decoder does, so the old-transformers sparse block contract works under UNSLOTH_FORCE_FLOAT32.

@danielhanchen

Copy link
Copy Markdown
Member Author

@codex review

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: ab2529b350

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)

if past_key_values is not None:
key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Pass cache metadata to Qwen cache updates

When fp16-forced Qwen3-MoE runs with use_cache on the still-supported 4.57.x transformers path, cache_position is received through **kwargs but is dropped here. The upstream attention forward passes {"sin": sin, "cos": cos, "cache_position": cache_position} into past_key_values.update; without that metadata, static caches cannot write the current token into the correct slot, so cached generation can fail or reuse corrupted K/V states. Forward the cache kwargs when they are present before calling update.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed: the update call now forwards cache_kwargs with sin, cos and cache_position, which 4.57.x static caches need; 5.x layer caches absorb the extra dict.

Comment on lines +170 to +173
if attn_impl == "flex_attention":
# Flex builds a BlockMask that raw SDPA cannot consume; dispatch to the
# interface (returns (b, q, h, d), already transposed).
attention_interface = ALL_ATTENTION_FUNCTIONS[attn_impl]

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Convert FlashAttention masks before raw SDPA

If the model is loaded with _attn_implementation='flash_attention_2', the model prepares a FlashAttention-style 2-D padding mask, but this patch falls through to raw scaled_dot_product_attention for every non-flex implementation. A 2-D (batch, kv_len) mask is not generally broadcastable to SDPA's (batch, heads, q_len, kv_len) scores, so batched fp16-forced Qwen3-MoE training/generation with padding can crash or mask the wrong dimension. Either dispatch non-SDPA implementations through ALL_ATTENTION_FUNCTIONS or expand the mask to SDPA's expected 4-D form before the call.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed: non-SDPA backends now dispatch through ALL_ATTENTION_FUNCTIONS with fp16 inputs (flash accumulates in float32 internally), so the 2-D padding mask keeps its native semantics.

Comment on lines +177 to +180
if attn_impl == "flex_attention":
# Flex builds a BlockMask that raw SDPA cannot consume; dispatch to the
# interface (returns (b, q, h, d), already transposed).
attention_interface = ALL_ATTENTION_FUNCTIONS[attn_impl]

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Convert FlashAttention masks before raw SDPA

The same mask-format issue exists for Gemma4: with _attn_implementation='flash_attention_2', upstream passes a 2-D padding mask suitable for the attention interface, but this branch sends it directly to scaled_dot_product_attention because only flex attention is dispatched. For batch sizes or sequence lengths where that mask cannot broadcast to (batch, heads, q_len, kv_len), fp16-forced Gemma4 runs with padding will fail or apply an incorrect mask. Convert the mask to SDPA format or keep using the registered attention interface for non-SDPA backends.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed the same way for Gemma4: registered backends get dispatched through the interface with fp16 inputs; raw SDPA is now only used for sdpa itself or unknown backends.

attn_output = attn_output.contiguous().reshape(*input_shape, -1)
attn_output = attn_output.to(torch.float16)
attn_output = self.o_proj(attn_output)
return attn_output, None

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P3 Badge Preserve requested attention weights

When callers set output_attentions=True or choose eager attention, upstream Qwen3-MoE can return attention weights, but this forced-fp32 replacement always returns None. Attention-distillation, diagnostics, or tests that rely on those weights will silently lose them only on the fp16-forced path; route requested eager attention through the attention interface or compute and return weights when requested.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added an eager branch that computes attention in float32 and returns the weights, restoring the module-level eager contract. Note the qwen3_moe decoder discards the weights on both 4.57.x and 5.x, so model-level outputs were never affected.

attn_output = attn_output.contiguous().reshape(*input_shape, -1)
attn_output = attn_output.to(torch.float16)
attn_output = self.o_proj(attn_output)
return attn_output, None

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P3 Badge Preserve requested attention weights

For Gemma4, callers can request attention tensors through output_attentions/eager attention, but this patched forward discards them by unconditionally returning None. That breaks any fp16-forced training or evaluation path that consumes attention maps for logging, distillation, or validation, while the non-forced path still returns them; preserve the upstream behavior when attentions are requested.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same eager branch added for Gemma4; weights are computed in float32 and returned. Gemma4 upstream has no output_attentions collection, so this only affects direct module calls.

Comment thread unsloth_zoo/model_lists.py Outdated
Comment on lines +49 to +50
"gemma4", # substring-matches "gemma4" and "gemma4_text"; fp16 NaNs the grad_norm in the backward (targeted float32 via gemma4_float32.py)
"glm4_moe", # GLM-4.x MoE (glm4_moe / glm4_moe_lite): fp16 NaNs grad_norm; bf16 load is dtype-consistent

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P3 Badge Add explicit FORCE_FLOAT32 variants for MLX

On the MLX loader path, _is_force_float32_arch does an exact normalized comparison rather than the CUDA substring match described above this list, so these entries match gemma4 and glm4_moe but not common variant model types like gemma4_text or glm4_moe_lite. Those variants will still downcast bf16 weights to fp16 without the warning this shared list is meant to trigger; add explicit normalized entries or update the MLX helper if the variants are intended to be covered.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added gemma4text and glm4_moe_lite entries so the MLX exact-match helper covers those config model_types; the CUDA substring path already matched them via the base entries.

@danielhanchen

Copy link
Copy Markdown
Member Author

@codex review

@chatgpt-codex-connector

Copy link
Copy Markdown

Codex Review: Didn't find any major issues. 🚀

Reviewed commit: ce3924b3db

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

@danielhanchen

Copy link
Copy Markdown
Member Author

@codex review

@chatgpt-codex-connector

Copy link
Copy Markdown

Codex Review: Didn't find any major issues. 🎉

Reviewed commit: ce3924b3db

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

@unslothai unslothai deleted a comment from chatgpt-codex-connector Bot Jul 6, 2026
@unslothai unslothai deleted a comment from chatgpt-codex-connector Bot Jul 6, 2026
@unslothai unslothai deleted a comment from chatgpt-codex-connector Bot Jul 6, 2026
@danielhanchen

Copy link
Copy Markdown
Member Author

float16 before / after

Quick summary of what this PR fixes. With dtype=torch.float16 these three MoE models used to keep a finite loss but NaN the grad_norm (mostly on the very first step, gemma4 most often), so the optimizer skipped the step and training limped along unreliably. After the fix they train finite and behave like the bf16 runs.

model fp16 grad_norm before fp16 grad_norm after note
gemma4 NaN on step ~1 (the most frequent of the three) finite over 15 steps, no NaN; loss tracks bf16 almost exactly (step 1 0.899 vs 0.901, step 15 0.628 vs 0.619) uses QK-norm, needs the targeted float32 patch
qwen3_moe (Qwen3-30B-A3B) NaN grad_norm, usually the first step finite over 15 steps, grad_norm 0.8-1.9 uses QK-norm, needs the targeted float32 patch
glm4_moe (GLM-4.x MoE) NaN grad_norm, usually the first step finite over 15 steps, grad_norm 0.17-0.38 dtype-consistent, the bf16 load alone trains finite

Root cause: a backward gradient overflow along the residual stream. Forward activations peak around 350, far below fp16's 65504, but the backward pass overflows, which is why the loss stays finite while the grad_norm goes NaN.

The same runs were always fully finite in bf16, which is what the fix leans on: a float16 request now loads bf16 weights and sets do_forced_float32, with targeted per-module float32 patches on the residual/norm/attention-core only for the two QK-norm models. Everything gates on UNSLOTH_FORCE_FLOAT32 == "1", so the bf16 path is untouched.

Note on the numbers: the "after" grad_norm ranges and the gemma4 loss pairs are measured (single B200, 4bit QLoRA, transformers 5.5, 15 steps); the "before" column is qualitative, since the pre-fix failure mode was a NaN grad_norm rather than a stable value to tabulate.

These MoE architectures silently NaN the grad_norm when loaded with
dtype=torch.float16 (a backward gradient overflow along the residual stream;
forward activations peak ~350, far below fp16's 65504, and bf16 is fully finite).
They are now handled the same way as gpt_oss / qwen3_5 / gemma3.

- model_lists.py: add gemma4, glm4_moe and qwen3_moe to FORCE_FLOAT32 so a
  float16 request loads bf16 weights and sets do_forced_float32.
- glm4_moe needs nothing further: its modules are dtype-consistent, so the bf16
  load alone trains finite.
- gemma4 and qwen3_moe have QK-norm, so do_forced_float32's generic fp32 upcast
  puts q/k in float32 while v stays fp16 and SDPA raises a dtype mismatch. Add
  targeted per-module float32 patches (gemma4_float32.py, qwen3_moe_float32.py),
  gated on UNSLOTH_FORCE_FLOAT32, that keep only the residual stream, the RMSNorm
  internals and the attention core (q/k/v/scores) in float32 while every
  projection, expert and router stays fp16. This mirrors the existing gemma3
  patches and keeps the fast path fast (no whole-model float32).

Validated on B200 (4bit QLoRA, transformers 5.5): float16 requests train with
finite loss and grad_norm over 15 steps for all three, bf16 is unchanged (the
patches no-op when UNSLOTH_FORCE_FLOAT32 is 0), and the float32 upcast is
confirmed targeted (heavy matmuls/experts stay fp16/4bit).
…port test

The submodule completeness check in test_temporary_patches_imports.py enforces
that every temporary_patches submodule on disk is listed in
TEMPORARY_PATCHES_SUBMODULES. The two new float32 patch modules were not added,
so the test failed. Add them so the import smoke coverage stays complete.
…at32 patches

The Gemma4 attention replacement swallowed the native shared_kv_states
dict that transformers 5.5.2+ passes, so kv-shared layers lost their
producer state; it now reads and writes the dict keyed by layer_type,
keeping the 5.5.0-5.5.1 carrier path as fallback. Both float32 attention
replacements dispatch flex_attention masks through the attention
interface instead of raw SDPA, mirroring the Gemma3 patch. The Qwen3-MoE
decoder replacement unpacks tuple outputs from old-transformers sparse
MoE blocks like the upstream decoder does.
… patches

Forward cache_kwargs (sin, cos, cache_position) to past_key_values.update
so 4.57.x static caches write the correct slot; 5.x layer caches absorb
the extra dict. Dispatch flash and other registered backends through
ALL_ATTENTION_FUNCTIONS with fp16 inputs (flash kernels accumulate in
float32 internally) so the 2-D padding mask keeps its native semantics,
and keep flex on the fp32 interface path. Compute eager attention in
float32 and return its weights to preserve the module-level eager
contract. Add gemma4text and glm4_moe_lite entries so the MLX
exact-match helper covers those variants.
@danielhanchen
danielhanchen force-pushed the fp16-moe-force-float32 branch from ce3924b to b222d42 Compare July 6, 2026 12:36
@danielhanchen
danielhanchen merged commit f7ad1e0 into main Jul 6, 2026
@chatgpt-codex-connector

Copy link
Copy Markdown

Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits.
Repo admins can enable using credits for code reviews in their settings.

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