diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index fc91178d88..a42b7f38a7 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -109,6 +109,13 @@ FORCE_FLOAT32 = [ "qwen3_5", # Qwen3.5 GDN layers produce NaN grad norms in float16 training ] +# Models that must use bfloat16 instead of float16. +# torch.compile backward graphs overflow fp16 intermediates for these models. +FORCE_BFLOAT16 = [ + "gemma4,", # Add comma bc gemma4 will match gemma4_text + "gemma4text", # Gemma4TextModel (standalone text-only Gemma4) +] + global DISABLE_COMPILE_MODEL_NAMES # Must be alphabetically sorted for each entry @@ -1373,6 +1380,24 @@ class FastModel(FastBaseModel): os.environ["UNSLOTH_FORCE_FLOAT32"] = "1" dtype = torch.bfloat16 # Change to bfloat16 loading break + # Switch fp16 to bf16 for models whose torch.compile backward overflows fp16 + global FORCE_BFLOAT16 + for disable_name in FORCE_BFLOAT16: + if ( + ( + disable_name.lower() + == model_type_arch.lower().replace("-", "").replace("_", "") + or disable_name.lower() in model_types_all + ) + and dtype == torch.float16 + and SUPPORTS_BFLOAT16 + ): + logger.warning_once( + f"Unsloth: {model_type_arch} does not support float16 training. " + f"Switching to bfloat16." + ) + dtype = torch.bfloat16 + break # Apply gradient checkpointing with smart heuristics use_gradient_checkpointing = apply_unsloth_gradient_checkpointing( use_gradient_checkpointing, max_seq_length, dtype