diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index cac5acd838..3cd8508ffa 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -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, ) diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index f05cc95d60..d212c224be 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -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,