From 739aa923fbd545c310cb719cec2b8311c58ab775 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 5 Jan 2026 06:56:57 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/sentence_transformer.py | 97 ++++++++++++++++++++------ 1 file changed, 76 insertions(+), 21 deletions(-) diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 587a2e6e82..09236e5a37 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -659,7 +659,15 @@ class FastSentenceTransformer(FastModel): return modules, True # Encoder model types that benefit from native torch.compile instead of Unsloth patching - ENCODER_MODEL_TYPES = {"mpnet", "bert", "distilbert", "roberta", "xlm-roberta", "albert", "electra"} + ENCODER_MODEL_TYPES = { + "mpnet", + "bert", + "distilbert", + "roberta", + "xlm-roberta", + "albert", + "electra", + } @staticmethod def from_pretrained( @@ -754,10 +762,11 @@ class FastSentenceTransformer(FastModel): # NOTE: The old Unsloth path is BROKEN for encoder models with torch 2.9+ due to # conflicting @torch.compile and @torch.compiler.disable decorators. # Set UNSLOTH_COMPILE_DISABLE=1 to disable torch.compile and use the old path. - is_encoder_model = model_type.lower() in FastSentenceTransformer.ENCODER_MODEL_TYPES + is_encoder_model = ( + model_type.lower() in FastSentenceTransformer.ENCODER_MODEL_TYPES + ) use_fast_encoder = os.environ.get("UNSLOTH_COMPILE_DISABLE", "0") != "1" if use_fast_encoder and is_encoder_model: - # torch.compile mode: "reduce-overhead" is optimal for training compile_mode = "reduce-overhead" @@ -776,7 +785,9 @@ class FastSentenceTransformer(FastModel): supports_sdpa = False if config is not None: try: - model_class = _get_model_class(config, kwargs.get("auto_model", AutoModel)._model_mapping) + model_class = _get_model_class( + config, kwargs.get("auto_model", AutoModel)._model_mapping + ) supports_sdpa = getattr(model_class, "_supports_sdpa", False) except: pass @@ -791,13 +802,18 @@ class FastSentenceTransformer(FastModel): # Print optimization status sdpa_str = " + SDPA" if supports_sdpa else "" if load_in_4bit: - print(f"Unsloth: Using fast encoder path for {model_type} with 4-bit quantization{sdpa_str}") + print( + f"Unsloth: Using fast encoder path for {model_type} with 4-bit quantization{sdpa_str}" + ) else: - print(f"Unsloth: Using fast encoder path for {model_type} (torch.compile{sdpa_str})") + print( + f"Unsloth: Using fast encoder path for {model_type} (torch.compile{sdpa_str})" + ) # Handle 4-bit quantization via BitsAndBytesConfig if load_in_4bit: from transformers import BitsAndBytesConfig + bnb_config = BitsAndBytesConfig( load_in_4bit = True, bnb_4bit_compute_dtype = dtype, @@ -811,7 +827,9 @@ class FastSentenceTransformer(FastModel): # Handle gradient checkpointing - warn user it conflicts with torch.compile _use_gc = use_gradient_checkpointing if _use_gc and _use_gc != False: - print("Unsloth Warning: Gradient checkpointing is incompatible with torch.compile.") + print( + "Unsloth Warning: Gradient checkpointing is incompatible with torch.compile." + ) print("Disabling torch.compile to enable gradient checkpointing.") compile_mode = None # Disable compilation @@ -850,7 +868,9 @@ class FastSentenceTransformer(FastModel): tokenizer.save_pretrained(save_directory) FastSentenceTransformer._add_unsloth_branding(save_directory) - st_model.save_pretrained_merged = types.MethodType(_save_pretrained_merged, st_model) + st_model.save_pretrained_merged = types.MethodType( + _save_pretrained_merged, st_model + ) def _push_to_hub_merged(self, repo_id, **push_kwargs): hub_token = push_kwargs.get("token", None) or get_token() @@ -858,22 +878,37 @@ class FastSentenceTransformer(FastModel): raise ValueError("No HF token provided") api = HfApi(token = hub_token) try: - api.create_repo(repo_id = repo_id, private = push_kwargs.get("private"), exist_ok = True, repo_type = "model") + api.create_repo( + repo_id = repo_id, + private = push_kwargs.get("private"), + exist_ok = True, + repo_type = "model", + ) except: pass FastSentenceTransformer._add_unsloth_tags(repo_id, hub_token) with tempfile.TemporaryDirectory() as temp_dir: self.save_pretrained_merged(temp_dir, **push_kwargs) - api.upload_folder(folder_path = temp_dir, repo_id = repo_id, commit_message = push_kwargs.get("commit_message", "Upload model")) + api.upload_folder( + folder_path = temp_dir, + repo_id = repo_id, + commit_message = push_kwargs.get( + "commit_message", "Upload model" + ), + ) print(f"Unsloth: Pushed to https://huggingface.co/{repo_id}") - st_model.push_to_hub_merged = types.MethodType(_push_to_hub_merged, st_model) + st_model.push_to_hub_merged = types.MethodType( + _push_to_hub_merged, st_model + ) return st_model # Warn if using 4-bit with encoder (slow due to dequantization overhead) if is_encoder_model and load_in_4bit: - print("Unsloth Warning: 4-bit quantization adds ~2.3x overhead for encoder models.") + print( + "Unsloth Warning: 4-bit quantization adds ~2.3x overhead for encoder models." + ) print("Consider using load_in_16bit=True for better performance.") # check if the model supports add_pooling_layer @@ -1111,8 +1146,11 @@ class FastSentenceTransformer(FastModel): inner_model = transformer_module.auto_model # Check if model is quantized (4-bit/8-bit) - is_quantized = getattr(inner_model, "is_quantized", False) or \ - getattr(inner_model.config, "quantization_config", None) is not None + is_quantized = ( + getattr(inner_model, "is_quantized", False) + or getattr(inner_model.config, "quantization_config", None) + is not None + ) # Track if gradient checkpointing was actually enabled gc_enabled = False @@ -1120,7 +1158,12 @@ class FastSentenceTransformer(FastModel): # Prepare for k-bit training if quantized if is_quantized: from peft import prepare_model_for_kbit_training - _gc_for_kbit = use_gradient_checkpointing if use_gradient_checkpointing else False + + _gc_for_kbit = ( + use_gradient_checkpointing + if use_gradient_checkpointing + else False + ) try: inner_model = prepare_model_for_kbit_training( inner_model, @@ -1131,12 +1174,16 @@ class FastSentenceTransformer(FastModel): except ValueError as e: if "does not support gradient checkpointing" in str(e): # Model doesn't support gradient checkpointing, disable it - print(f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping.") + print( + f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping." + ) inner_model = prepare_model_for_kbit_training( inner_model, use_gradient_checkpointing = False, ) - print("Unsloth: Prepared quantized model for k-bit training (without gradient checkpointing)") + print( + "Unsloth: Prepared quantized model for k-bit training (without gradient checkpointing)" + ) else: raise @@ -1149,7 +1196,9 @@ class FastSentenceTransformer(FastModel): gc_enabled = True except ValueError as e: if "does not support gradient checkpointing" in str(e): - print(f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping.") + print( + f"Unsloth Warning: {inner_model.__class__.__name__} does not support gradient checkpointing. Skipping." + ) # Create LoRA config lora_config = LoraConfig( @@ -1169,13 +1218,19 @@ class FastSentenceTransformer(FastModel): # Re-enable torch.compile if gradient checkpointing was requested but couldn't be enabled if compile_mode is None and not gc_enabled: compile_mode = "reduce-overhead" - print("Unsloth: Re-enabling torch.compile since gradient checkpointing is not supported") + print( + "Unsloth: Re-enabling torch.compile since gradient checkpointing is not supported" + ) if compile_mode is not None: - print(f"Unsloth: Applying torch.compile with mode='{compile_mode}' for 6x speedup") + print( + f"Unsloth: Applying torch.compile with mode='{compile_mode}' for 6x speedup" + ) peft_model = torch.compile(peft_model, mode = compile_mode) else: - print("Unsloth: torch.compile disabled (gradient checkpointing enabled)") + print( + "Unsloth: torch.compile disabled (gradient checkpointing enabled)" + ) # Re-assign the peft model back to the transformer module transformer_module.auto_model = peft_model