[pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci
This commit is contained in:
pre-commit-ci[bot] 2026-03-12 10:01:35 +00:00
commit b85e7d0488
3 changed files with 20 additions and 7 deletions

View file

@ -226,10 +226,16 @@ def determine_attention_implementation(model_class, config):
model_type = getattr(config, "model_type", "").lower()
# 1. Flash Attention 2
if HAS_FLASH_ATTENTION and model_type not in ("gpt_oss", "mllama") and not model_type.startswith("gemma3"):
if (
HAS_FLASH_ATTENTION
and model_type not in ("gpt_oss", "mllama")
and not model_type.startswith("gemma3")
):
supports_fa2 = False
if model_class is not None:
supports_fa2 = getattr(model_class, "_supports_flash_attn_2", False) or getattr(model_class, "_supports_flash_attn", False)
supports_fa2 = getattr(
model_class, "_supports_flash_attn_2", False
) or getattr(model_class, "_supports_flash_attn", False)
if supports_fa2:
if config is not None:
@ -243,8 +249,10 @@ def determine_attention_implementation(model_class, config):
try:
from transformers.utils.import_utils import is_torch_flex_attn_available
if is_torch_flex_attn_available() and (model_class is not None) and getattr(
model_class, "_supports_flex_attn", False
if (
is_torch_flex_attn_available()
and (model_class is not None)
and getattr(model_class, "_supports_flex_attn", False)
):
# GPT-OSS, Mllama and Gemma3 use eager/sdpa attention during
# inference since flex attention returns incorrect results or errors out.
@ -254,7 +262,10 @@ def determine_attention_implementation(model_class, config):
# 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("gemma3"):
if model_type not in (
"gpt_oss",
"mllama",
) and not model_type.startswith("gemma3"):
if config is not None:
setattr(config, "_attn_implementation", "flex_attention")
if hasattr(config, "attn_implementation"):

View file

@ -2341,7 +2341,9 @@ class FastLlamaModel:
model_function = MODEL_FOR_CAUSAL_LM_MAPPING[model_config.__class__]
IS_FALCON_H1 = model_config.model_type.startswith("falcon_h1")
preferred_attn_impl = determine_attention_implementation(model_function, model_config)
preferred_attn_impl = determine_attention_implementation(
model_function, model_config
)
has_rope_scaling = False
try:

View file

@ -608,7 +608,7 @@ class FastBaseModel:
model_class = auto_model._model_mapping[auto_config.__class__]
except Exception:
model_class = None
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