diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index d47762f0b0..a97bd0eef6 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -79,63 +79,88 @@ class FastSentenceTransformer(FastModel): return mode except Exception as e: - print( - f"Failed to detect pooling mode: {e}, defaulting to mean pooling." - ) + print(f"Failed to detect pooling mode: {e}, defaulting to mean pooling.") return "mean" - + @staticmethod - def _load_modules(model_name, token, model, tokenizer, max_seq_length, pooling_mode, trust_remote_code = False): + def _load_modules( + model_name, + token, + model, + tokenizer, + max_seq_length, + pooling_mode, + trust_remote_code = False, + ): modules = OrderedDict() - + # grope around for modules.json modules_json_path = None - 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") + 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: - try: - modules_json_path = hf_hub_download(model_name, "modules.json", token=token) - except: - pass + try: + modules_json_path = hf_hub_download( + model_name, "modules.json", token = token + ) + except: + pass if modules_json_path and os.path.exists(modules_json_path): - with open(modules_json_path, encoding="utf8") as f: + with open(modules_json_path, encoding = "utf8") as f: modules_config = json.load(f) for module_config in modules_config: class_ref = module_config["type"] - name = module_config["name"] if "name" in module_config else str(module_config.get("idx", len(modules))) - + name = ( + module_config["name"] + if "name" in module_config + else str(module_config.get("idx", len(modules))) + ) + # main module if class_ref == "sentence_transformers.models.Transformer": transformer_module = Transformer( - model_name, - max_seq_length=max_seq_length, - model_args = {"trust_remote_code" : trust_remote_code}, - config_args = {"trust_remote_code" : trust_remote_code}, + model_name, + max_seq_length = max_seq_length, + model_args = {"trust_remote_code": trust_remote_code}, + config_args = {"trust_remote_code": trust_remote_code}, ) transformer_module.auto_model = model transformer_module.tokenizer = tokenizer - + # move tokenizer do_lower_case to transformer module - transformer_module.do_lower_case = getattr(tokenizer, "do_lower_case", False) - model_forward_params = list(inspect.signature(model.forward).parameters) - transformer_module.model_forward_params = set(model_forward_params) | { - "input_ids", "attention_mask", "token_type_ids", "inputs_embeds", + transformer_module.do_lower_case = getattr( + tokenizer, "do_lower_case", False + ) + model_forward_params = list( + inspect.signature(model.forward).parameters + ) + transformer_module.model_forward_params = set( + model_forward_params + ) | { + "input_ids", + "attention_mask", + "token_type_ids", + "inputs_embeds", } if max_seq_length is None: - pass + pass # is this overkill? should we just force user to set it? current_max_seq = max_seq_length if current_max_seq is None: - if hasattr(model, "config") and hasattr(model.config, "max_position_embeddings"): - current_max_seq = model.config.max_position_embeddings - elif hasattr(tokenizer, "model_max_length"): - current_max_seq = tokenizer.model_max_length - else: - current_max_seq = 512 - + if hasattr(model, "config") and hasattr( + model.config, "max_position_embeddings" + ): + current_max_seq = model.config.max_position_embeddings + elif hasattr(tokenizer, "model_max_length"): + current_max_seq = tokenizer.model_max_length + else: + current_max_seq = 512 + transformer_module.max_seq_length = current_max_seq transformer_module.config_keys = ["max_seq_length", "do_lower_case"] transformer_module.save_in_root = True @@ -143,54 +168,69 @@ class FastSentenceTransformer(FastModel): model.config.tokenizer_class = tokenizer.__class__.__name__ modules[name] = transformer_module - + # load other modules else: module_path = module_config["path"] if os.path.isdir(model_name): - load_path = os.path.join(model_name, module_path) + load_path = os.path.join(model_name, module_path) else: - # still looking - try: - load_path = load_dir_path(model_name, module_path, token=token) - except: - print(f"Unsloth Warning: Could not download module {module_path} for {class_ref}. Skipping.") - continue - + # still looking + try: + load_path = load_dir_path( + model_name, module_path, token = token + ) + except: + print( + f"Unsloth Warning: Could not download module {module_path} for {class_ref}. Skipping." + ) + continue + module_class = import_from_string(class_ref) # load module try: module = module_class.load(load_path) modules[name] = module except Exception as e: - print(f"Unsloth Warning: Failed to load module {name} ({class_ref}) from {load_path}: {e}") + print( + f"Unsloth Warning: Failed to load module {name} ({class_ref}) from {load_path}: {e}" + ) else: # fallback if no modules.json, is this necessary? - print("Unsloth: No modules.json found, falling back to [Transformer, Pooling, Normalize]") + print( + "Unsloth: No modules.json found, falling back to [Transformer, Pooling, Normalize]" + ) transformer_module = Transformer( - model_name, - max_seq_length=max_seq_length, - model_args = {"trust_remote_code" : trust_remote_code}, - config_args = {"trust_remote_code" : trust_remote_code}, + model_name, + max_seq_length = max_seq_length, + model_args = {"trust_remote_code": trust_remote_code}, + config_args = {"trust_remote_code": trust_remote_code}, ) transformer_module.auto_model = model transformer_module.tokenizer = tokenizer - + # move tokenizer do_lower_case to transformer module - transformer_module.do_lower_case = getattr(tokenizer, "do_lower_case", False) + transformer_module.do_lower_case = getattr( + tokenizer, "do_lower_case", False + ) model_forward_params = list(inspect.signature(model.forward).parameters) transformer_module.model_forward_params = set(model_forward_params) | { - "input_ids", "attention_mask", "token_type_ids", "inputs_embeds", + "input_ids", + "attention_mask", + "token_type_ids", + "inputs_embeds", } - + if max_seq_length is None: - if hasattr(model, "config") and hasattr(model.config, "max_position_embeddings"): - max_seq_length = model.config.max_position_embeddings + if hasattr(model, "config") and hasattr( + model.config, "max_position_embeddings" + ): + max_seq_length = model.config.max_position_embeddings elif hasattr(tokenizer, "model_max_length"): - max_seq_length = tokenizer.model_max_length + max_seq_length = tokenizer.model_max_length else: - max_seq_length = 512 + max_seq_length = 512 transformer_module.max_seq_length = max_seq_length transformer_module.config_keys = ["max_seq_length", "do_lower_case"] transformer_module.save_in_root = True @@ -199,15 +239,21 @@ class FastSentenceTransformer(FastModel): model.config.tokenizer_class = tokenizer.__class__.__name__ modules["0"] = transformer_module - - hidden_size = model.config.hidden_size if hasattr(model.config, "hidden_size") else 768 - + + hidden_size = ( + model.config.hidden_size + if hasattr(model.config, "hidden_size") + else 768 + ) + if pooling_mode == "mean": - pooling_mode = FastSentenceTransformer._read_pooling_mode(model_name, token) + pooling_mode = FastSentenceTransformer._read_pooling_mode( + model_name, token + ) pooling_module = Pooling( - word_embedding_dimension=hidden_size, - pooling_mode=pooling_mode, + word_embedding_dimension = hidden_size, + pooling_mode = pooling_mode, ) # end of fallback modules["1"] = pooling_module @@ -266,7 +312,7 @@ class FastSentenceTransformer(FastModel): # however unsloth throws an exception if "UNSLOTH_WARN_UNINITIALIZED" == 1 and it sees unused weights old_environ = os.environ.get("UNSLOTH_WARN_UNINITIALIZED", "1") os.environ["UNSLOTH_WARN_UNINITIALIZED"] = "0" - + try: model, tokenizer = FastModel.from_pretrained( model_name = model_name, @@ -300,14 +346,23 @@ class FastSentenceTransformer(FastModel): # try to load modules, otherwise fallback to old hard-coded modules from sentence_transformers import SentenceTransformer - modules = FastSentenceTransformer._load_modules(model_name, token, model, tokenizer, max_seq_length, pooling_mode, trust_remote_code=trust_remote_code) - st_model = SentenceTransformer(modules=modules, device=device_map) + modules = FastSentenceTransformer._load_modules( + model_name, + token, + model, + tokenizer, + max_seq_length, + pooling_mode, + trust_remote_code = trust_remote_code, + ) + + st_model = SentenceTransformer(modules = modules, device = device_map) def _save_pretrained_merged(self, save_directory, **kwargs): # sentence-transformers config and modules only get saved if we call save_pretrained self.save_pretrained(save_directory) - + # remove LoRA adapters since we are saving the merged model for file in ["adapter_model.safetensors", "adapter_config.json"]: try: @@ -317,9 +372,13 @@ class FastSentenceTransformer(FastModel): # save merged weights tokenizer = kwargs.pop("tokenizer", self.tokenizer) - self[0].auto_model.save_pretrained_merged(save_directory, tokenizer=tokenizer, **kwargs) + self[0].auto_model.save_pretrained_merged( + save_directory, tokenizer = tokenizer, **kwargs + ) - st_model.save_pretrained_merged = types.MethodType(_save_pretrained_merged, st_model) + st_model.save_pretrained_merged = types.MethodType( + _save_pretrained_merged, st_model + ) return st_model @staticmethod