From 855f719e7667091bd854e2579e7060c1186bdf64 Mon Sep 17 00:00:00 2001 From: electroglyph Date: Fri, 12 Dec 2025 03:45:13 -0800 Subject: [PATCH] Gemini code review suggestions --- unsloth/models/sentence_transformer.py | 243 ++++++++++++------------- 1 file changed, 112 insertions(+), 131 deletions(-) diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index c6b0fbc4c0..3c3595e20c 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -13,39 +13,42 @@ # limitations under the License. from .loader import FastModel - +import torch +import inspect +import json +import os +from huggingface_hub import hf_hub_download class FastSentenceTransformer(FastModel): @staticmethod def from_pretrained( model_name, - max_seq_length = None, - dtype = None, - load_in_4bit = True, - load_in_8bit = False, - load_in_16bit = False, - full_finetuning = False, - token = None, - device_map = "sequential", - rope_scaling = None, - fix_tokenizer = True, - trust_remote_code = False, - use_gradient_checkpointing = "unsloth", - resize_model_vocab = None, - revision = None, - use_exact_model_name = False, - offload_embedding = False, - random_state = 3407, - max_lora_rank = 64, - disable_log_stats = True, - qat_scheme = None, - load_in_fp8 = False, - unsloth_tiled_mlp = False, - pooling_mode = "mean", + max_seq_length=None, + dtype=None, + load_in_4bit=True, + load_in_8bit=False, + load_in_16bit=False, + full_finetuning=False, + token=None, + device_map="sequential", + rope_scaling=None, + fix_tokenizer=True, + trust_remote_code=False, + use_gradient_checkpointing="unsloth", + resize_model_vocab=None, + revision=None, + use_exact_model_name=False, + offload_embedding=False, + random_state=3407, + max_lora_rank=64, + disable_log_stats=True, + qat_scheme=None, + load_in_fp8=False, + unsloth_tiled_mlp=False, + pooling_mode="mean", **kwargs, ): try: - import sentence_transformers from sentence_transformers import SentenceTransformer from sentence_transformers.models import Transformer, Pooling, Normalize from transformers import AutoModel @@ -62,46 +65,41 @@ class FastSentenceTransformer(FastModel): kwargs["add_pooling_layer"] = False model, tokenizer = FastModel.from_pretrained( - model_name = model_name, - max_seq_length = max_seq_length, - dtype = dtype, - load_in_4bit = load_in_4bit, - load_in_8bit = load_in_8bit, - load_in_16bit = load_in_16bit, - full_finetuning = full_finetuning, - token = token, - device_map = device_map, - rope_scaling = rope_scaling, - fix_tokenizer = fix_tokenizer, - trust_remote_code = trust_remote_code, - use_gradient_checkpointing = use_gradient_checkpointing, - resize_model_vocab = resize_model_vocab, - revision = revision, - return_logits = False, - use_exact_model_name = use_exact_model_name, - offload_embedding = offload_embedding, - random_state = random_state, - max_lora_rank = max_lora_rank, - disable_log_stats = disable_log_stats, - qat_scheme = qat_scheme, - load_in_fp8 = load_in_fp8, - unsloth_tiled_mlp = unsloth_tiled_mlp, + model_name=model_name, + max_seq_length=max_seq_length, + dtype=dtype, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + load_in_16bit=load_in_16bit, + full_finetuning=full_finetuning, + token=token, + device_map=device_map, + rope_scaling=rope_scaling, + fix_tokenizer=fix_tokenizer, + trust_remote_code=trust_remote_code, + use_gradient_checkpointing=use_gradient_checkpointing, + resize_model_vocab=resize_model_vocab, + revision=revision, + return_logits=False, + use_exact_model_name=use_exact_model_name, + offload_embedding=offload_embedding, + random_state=random_state, + max_lora_rank=max_lora_rank, + disable_log_stats=disable_log_stats, + qat_scheme=qat_scheme, + load_in_fp8=load_in_fp8, + unsloth_tiled_mlp=unsloth_tiled_mlp, **kwargs, ) transformer_module = Transformer.__new__(Transformer) - import torch - torch.nn.Module.__init__(transformer_module) - transformer_module.auto_model = model transformer_module.tokenizer = tokenizer transformer_module.do_lower_case = False if hasattr(tokenizer, "do_lower_case"): transformer_module.do_lower_case = tokenizer.do_lower_case - import inspect - model_forward_params = list(inspect.signature(model.forward).parameters) transformer_module.model_forward_params = set(model_forward_params) | { "input_ids", @@ -137,17 +135,13 @@ class FastSentenceTransformer(FastModel): # detect pooling mode if not specified/default if pooling_mode == "mean": try: - from huggingface_hub import hf_hub_download - import json - import os - if os.path.exists(model_name) and os.path.exists( os.path.join(model_name, "modules.json") ): modules_json_path = os.path.join(model_name, "modules.json") else: modules_json_path = hf_hub_download( - model_name, "modules.json", token = token + model_name, "modules.json", token=token ) with open(modules_json_path, "r") as f: @@ -169,37 +163,24 @@ class FastSentenceTransformer(FastModel): pooling_config_path = hf_hub_download( model_name, os.path.join(pooling_path, "config.json"), - token = token, + token=token, ) break if pooling_config_path: with open(pooling_config_path, "r") as f: pooling_config = json.load(f) - if ( - "pooling_mode_cls_token" in pooling_config - and pooling_config["pooling_mode_cls_token"] - ): - print("Pooling mode detected as cls, updating...") - pooling_mode = "cls" - elif ( - "pooling_mode_mean_tokens" in pooling_config - and pooling_config["pooling_mode_mean_tokens"] - ): - print("Pooling mode detected as mean, updating...") - pooling_mode = "mean" - elif ( - "pooling_mode_max_tokens" in pooling_config - and pooling_config["pooling_mode_max_tokens"] - ): - print("Pooling mode detected as max, updating...") - pooling_mode = "max" - elif ( - "pooling_mode_mean_sqrt_len_tokens" in pooling_config - and pooling_config["pooling_mode_mean_sqrt_len_tokens"] - ): - print("Pooling mode detected as mean_sqrt_len, updating...") - pooling_mode = "mean_sqrt_len" + pooling_map = { + "pooling_mode_cls_token": "cls", + "pooling_mode_mean_tokens": "mean", + "pooling_mode_max_tokens": "max", + "pooling_mode_mean_sqrt_len_tokens": "mean_sqrt_len", + } + for config_key, mode in pooling_map.items(): + if pooling_config.get(config_key): + print(f"Pooling mode detected as {mode}, updating...") + pooling_mode = mode + break except Exception as e: print( @@ -207,36 +188,36 @@ class FastSentenceTransformer(FastModel): ) pooling_module = Pooling( - word_embedding_dimension = hidden_size, - pooling_mode = pooling_mode, + word_embedding_dimension=hidden_size, + pooling_mode=pooling_mode, ) normalize_module = Normalize() modules = [transformer_module, pooling_module, normalize_module] - st_model = SentenceTransformer(modules = modules) + st_model = SentenceTransformer(modules=modules) return st_model @staticmethod def get_peft_model( model, - r = 16, - target_modules = [ + r=16, + target_modules=[ "query", "key", "value", "dense", ], - lora_alpha = 16, - lora_dropout = 0.0, - bias = "none", - layers_to_transform = None, - layers_pattern = None, - use_gradient_checkpointing = "unsloth", - random_state = 3407, - max_seq_length = 2048, - use_rslora = False, - modules_to_save = None, - init_lora_weights = True, - loftq_config = {}, + lora_alpha=16, + lora_dropout=0.0, + bias="none", + layers_to_transform=None, + layers_pattern=None, + use_gradient_checkpointing="unsloth", + random_state=3407, + max_seq_length=2048, + use_rslora=False, + modules_to_save=None, + init_lora_weights=True, + loftq_config={}, **kwargs, ): from sentence_transformers import SentenceTransformer @@ -251,21 +232,21 @@ class FastSentenceTransformer(FastModel): inner_model = transformer_module.auto_model peft_model = FastModel.get_peft_model( - model = inner_model, - r = r, - target_modules = target_modules, - lora_alpha = lora_alpha, - lora_dropout = lora_dropout, - bias = bias, - layers_to_transform = layers_to_transform, - layers_pattern = layers_pattern, - use_gradient_checkpointing = use_gradient_checkpointing, - random_state = random_state, - max_seq_length = max_seq_length, - use_rslora = use_rslora, - modules_to_save = modules_to_save, - init_lora_weights = init_lora_weights, - loftq_config = loftq_config, + model=inner_model, + r=r, + target_modules=target_modules, + lora_alpha=lora_alpha, + lora_dropout=lora_dropout, + bias=bias, + layers_to_transform=layers_to_transform, + layers_pattern=layers_pattern, + use_gradient_checkpointing=use_gradient_checkpointing, + random_state=random_state, + max_seq_length=max_seq_length, + use_rslora=use_rslora, + modules_to_save=modules_to_save, + init_lora_weights=init_lora_weights, + loftq_config=loftq_config, **kwargs, ) @@ -274,20 +255,20 @@ class FastSentenceTransformer(FastModel): return model else: return FastModel.get_peft_model( - model = model, - r = r, - target_modules = target_modules, - lora_alpha = lora_alpha, - lora_dropout = lora_dropout, - bias = bias, - layers_to_transform = layers_to_transform, - layers_pattern = layers_pattern, - use_gradient_checkpointing = use_gradient_checkpointing, - random_state = random_state, - max_seq_length = max_seq_length, - use_rslora = use_rslora, - modules_to_save = modules_to_save, - init_lora_weights = init_lora_weights, - loftq_config = loftq_config, + model=model, + r=r, + target_modules=target_modules, + lora_alpha=lora_alpha, + lora_dropout=lora_dropout, + bias=bias, + layers_to_transform=layers_to_transform, + layers_pattern=layers_pattern, + use_gradient_checkpointing=use_gradient_checkpointing, + random_state=random_state, + max_seq_length=max_seq_length, + use_rslora=use_rslora, + modules_to_save=modules_to_save, + init_lora_weights=init_lora_weights, + loftq_config=loftq_config, **kwargs, )