feat: Support custom auto_model for wider model compatibility (Whisper, Bert,etc) & attn_implementation support (#2263)

* Update loader.py

* Update vision.py

* Update vision.py

fix attn_implementation

* Refactor: Improve parameter handling and checks in loader/vision
This commit is contained in:
Etherll 2025-04-14 23:10:05 +02:00 committed by GitHub
commit 68d1c230ab
2 changed files with 35 additions and 12 deletions

View file

@ -469,10 +469,14 @@ class FastModel(FastBaseModel):
return_logits = False, # Return logits
fullgraph = True, # No graph breaks
use_exact_model_name = False,
auto_model = None,
whisper_language = None,
whisper_task = None,
*args, **kwargs,
):
if token is None: token = get_token()
if whisper_language is not None: assert(type(whisper_language) is str)
if whisper_task is not None: assert(type(whisper_task) is str)
SUPPORTS_BFLOAT16 = is_bfloat16_supported()
if dtype is None:
dtype = torch.float16 if not SUPPORTS_BFLOAT16 else torch.bfloat16
@ -709,7 +713,8 @@ class FastModel(FastBaseModel):
# Check if VLM
is_vlm = any(x.endswith("ForConditionalGeneration") for x in model_config.architectures)
is_vlm = is_vlm or hasattr(model_config, "vision_config")
auto_model = AutoModelForVision2Seq if is_vlm else AutoModelForCausalLM
if auto_model is None:
auto_model = AutoModelForVision2Seq if is_vlm else AutoModelForCausalLM
model, tokenizer = FastBaseModel.from_pretrained(
model_name = model_name,
@ -727,6 +732,8 @@ class FastModel(FastBaseModel):
auto_model = auto_model,
use_gradient_checkpointing = use_gradient_checkpointing,
supports_sdpa = supports_sdpa,
whisper_language = whisper_language,
whisper_task = whisper_task,
*args, **kwargs,
)

View file

@ -236,6 +236,8 @@ class FastBaseModel:
auto_model = AutoModelForVision2Seq,
use_gradient_checkpointing = "unsloth",
supports_sdpa = True,
whisper_language = None,
whisper_task = None,
**kwargs,
):
if model_types is None:
@ -304,7 +306,8 @@ class FastBaseModel:
do_forced_float32 = True
pass
# Stop SDPA for some archs like Pixtral / Mistral3
kwargs["attn_implementation"] = "sdpa"
if not ("attn_implementation" in kwargs):
kwargs["attn_implementation"] = "sdpa"
if not supports_sdpa:
print(f"Unsloth: {model_type_arch.title()} does not support SDPA - switching to eager!")
del kwargs["attn_implementation"]
@ -352,6 +355,7 @@ class FastBaseModel:
# Check if using forced float32 - we load it in bfloat16, then cast to float16!
torch_dtype = dtype
if do_forced_float32: torch_dtype = torch.bfloat16
model = auto_model.from_pretrained(
model_name,
device_map = device_map,
@ -367,12 +371,23 @@ class FastBaseModel:
# Counteract saved tokenizers
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
auto_processor = AutoProcessor if auto_model is AutoModelForVision2Seq else AutoTokenizer
tokenizer = auto_processor.from_pretrained(
tokenizer_name,
padding_side = "right",
token = token,
)
is_vlm = (auto_model is AutoModelForVision2Seq)
is_whisper = (whisper_language is not None and whisper_task is not None)
auto_processor = AutoProcessor if (is_vlm or is_whisper) else AutoTokenizer
if whisper_language and whisper_task:
tokenizer = auto_processor.from_pretrained(
tokenizer_name,
padding_side = "right",
token = token,
language = whisper_language,
task = whisper_task,
)
else:
tokenizer = auto_processor.from_pretrained(
tokenizer_name,
padding_side = "right",
token = token,
)
if hasattr(tokenizer, "tokenizer"):
__tokenizer = tokenizer.tokenizer
# Add padding side as well
@ -469,6 +484,7 @@ class FastBaseModel:
modules_to_save = None,
init_lora_weights = True,
loftq_config = {},
task_type = TaskType.CAUSAL_LM,
temporary_location = "_unsloth_temporary_saved_buffers",
**kwargs,
):
@ -492,7 +508,7 @@ class FastBaseModel:
finetune_attention_modules = True
finetune_mlp_modules = True
pass
if target_modules is None:
if target_modules is None or target_modules == "all-linear":
target_modules = get_peft_regex(
model,
finetune_vision_layers = finetune_vision_layers,
@ -503,7 +519,7 @@ class FastBaseModel:
else:
assert(type(target_modules) in (list, tuple,))
pass
# Clear deleted GPU items
for _ in range(3):
gc.collect()
@ -516,7 +532,7 @@ class FastBaseModel:
target_modules = target_modules,
lora_dropout = lora_dropout,
bias = bias,
task_type = TaskType.CAUSAL_LM,
task_type = task_type,
)
model = prepare_model_for_kbit_training(
model,