Fix attention selection regressions in PR #4250

- Extract _set_attn_impl helper to reduce config-stamping boilerplate
- Add nemotron_h to FA2 and flex exclusion lists (from main)
- Add attention_dropout > 0 check for flex attention (from main)
- Restore gemma3n eager default in vision.py (timm vision towers)
- Preserve sdpa fallback for unmapped/remote-code vision configs
- Restore original SDPA warning default (print when env var unset)
- Make config=None safe in determine_attention_implementation
This commit is contained in:
Daniel Han 2026-03-31 16:47:48 +00:00
commit f77a42e0ab
2 changed files with 44 additions and 27 deletions

View file

@ -222,11 +222,24 @@ def apply_unsloth_gradient_checkpointing(
return use_gradient_checkpointing
def _set_attn_impl(config, impl):
"""Stamp the chosen attention implementation onto the config object."""
if config is not None:
setattr(config, "_attn_implementation", impl)
if hasattr(config, "attn_implementation"):
setattr(config, "attn_implementation", impl)
return impl
def determine_attention_implementation(model_class, config):
model_type = getattr(config, "model_type", "").lower()
model_type = getattr(config, "model_type", "").lower() if config else ""
# 1. Flash Attention 2
if HAS_FLASH_ATTENTION and model_type not in ("gpt_oss", "mllama") and not model_type.startswith("gemma3n"):
if (
HAS_FLASH_ATTENTION
and model_type not in ("gpt_oss", "mllama", "nemotron_h")
and not model_type.startswith("gemma3n")
):
supports_fa2 = False
if model_class is not None:
supports_fa2 = getattr(
@ -234,11 +247,7 @@ def determine_attention_implementation(model_class, config):
) or getattr(model_class, "_supports_flash_attn", False)
if supports_fa2:
if config is not None:
setattr(config, "_attn_implementation", "flash_attention_2")
if hasattr(config, "attn_implementation"):
setattr(config, "attn_implementation", "flash_attention_2")
return "flash_attention_2"
return _set_attn_impl(config, "flash_attention_2")
# 2. Flex Attention
if os.environ.get("UNSLOTH_ENABLE_FLEX_ATTENTION", "1") != "0":
@ -250,37 +259,33 @@ def determine_attention_implementation(model_class, config):
and (model_class is not None)
and getattr(model_class, "_supports_flex_attn", False)
):
# GPT-OSS, Mllama and Gemma3N use eager/sdpa attention during
# inference since flex attention returns incorrect results or errors out.
attention_dropout = getattr(config, "attention_dropout", 0) or 0
# GPT-OSS, Mllama, Gemma3N and NemotronH use eager/sdpa attention
# during inference since flex attention returns incorrect results
# or errors out.
# GPT-OSS: left padding issues cause incorrect outputs.
# Mllama: _update_causal_mask uses make_flex_block_causal_mask which
# creates BlockMask with Q_LEN=KV_LEN=total_seq_len, but during
# decode q_len=1, causing ValueError. Needs transformers update.
# Gemma3N: timm vision wrappers (eg Gemma3nVisionConfig) do not
# support flex_attention.
if model_type not in ("gpt_oss", "mllama") and not model_type.startswith("gemma3n"):
if config is not None:
setattr(config, "_attn_implementation", "flex_attention")
if hasattr(config, "attn_implementation"):
setattr(config, "attn_implementation", "flex_attention")
return "flex_attention"
# NemotronH: hybrid Mamba-2 + Transformer model that does not
# support flex_attention (raises NotImplementedError from transformers).
is_excluded = (
model_type in ("gpt_oss", "mllama", "nemotron_h")
or model_type.startswith("gemma3n")
)
if attention_dropout == 0 and not is_excluded:
return _set_attn_impl(config, "flex_attention")
except Exception:
pass
# 3. SDPA
if model_class is not None and getattr(model_class, "_supports_sdpa", False):
if config is not None:
setattr(config, "_attn_implementation", "sdpa")
if hasattr(config, "attn_implementation"):
setattr(config, "attn_implementation", "sdpa")
return "sdpa"
return _set_attn_impl(config, "sdpa")
# 4. Eager
if config is not None:
setattr(config, "_attn_implementation", "eager")
if hasattr(config, "attn_implementation"):
setattr(config, "attn_implementation", "eager")
return "eager"
return _set_attn_impl(config, "eager")
def _run_temporary_patches(phase):

View file

@ -609,7 +609,19 @@ class FastBaseModel:
except Exception:
model_class = None
attn_impl = determine_attention_implementation(model_class, auto_config)
model_type = str(getattr(auto_config, "model_type", "")).lower()
if model_type.startswith("gemma3n"):
# Gemma3N variants use timm-based vision towers which do not support
# flex_attention. The old code defaulted gemma3n to eager; preserve
# that behavior rather than letting the hierarchy pick sdpa.
attn_impl = "eager"
elif model_class is None and supports_sdpa:
# When model_class cannot be resolved (remote-code or unmapped
# configs), the old code defaulted to sdpa. Preserve that fallback
# instead of falling through to eager.
attn_impl = "sdpa"
else:
attn_impl = determine_attention_implementation(model_class, auto_config)
# Handle FP8 models: get_model_name has already redirected this to BF16 sibling if the model ships with
# FP8 weights. We just need to update it here for sanity.
@ -623,7 +635,7 @@ class FastBaseModel:
if not ("attn_implementation" in kwargs):
kwargs["attn_implementation"] = attn_impl
if not supports_sdpa and kwargs.get("attn_implementation") == "sdpa":
if os.environ.get("UNSLOTH_ENABLE_FLEX_ATTENTION", "1") == "0":
if os.environ.get("UNSLOTH_ENABLE_FLEX_ATTENTION", "0") == "0":
print(
f"Unsloth: {model_type_arch.title()} does not support SDPA - switching to fast eager."
)