Compare commits
3 commits
main
...
fix/gemma4
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e264d42413 | ||
|
|
e5ee16b8a1 | ||
|
|
d323b47805 |
1 changed files with 25 additions and 0 deletions
|
|
@ -109,6 +109,13 @@ FORCE_FLOAT32 = [
|
||||||
"qwen3_5", # Qwen3.5 GDN layers produce NaN grad norms in float16 training
|
"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
|
global DISABLE_COMPILE_MODEL_NAMES
|
||||||
# Must be alphabetically sorted for each entry
|
# Must be alphabetically sorted for each entry
|
||||||
|
|
||||||
|
|
@ -1373,6 +1380,24 @@ class FastModel(FastBaseModel):
|
||||||
os.environ["UNSLOTH_FORCE_FLOAT32"] = "1"
|
os.environ["UNSLOTH_FORCE_FLOAT32"] = "1"
|
||||||
dtype = torch.bfloat16 # Change to bfloat16 loading
|
dtype = torch.bfloat16 # Change to bfloat16 loading
|
||||||
break
|
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
|
# Apply gradient checkpointing with smart heuristics
|
||||||
use_gradient_checkpointing = apply_unsloth_gradient_checkpointing(
|
use_gradient_checkpointing = apply_unsloth_gradient_checkpointing(
|
||||||
use_gradient_checkpointing, max_seq_length, dtype
|
use_gradient_checkpointing, max_seq_length, dtype
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue