diff --git a/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py b/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py new file mode 100644 index 0000000000..ff9b91ec23 --- /dev/null +++ b/tests/python/test_fast_sentence_transformer_redirect_lifecycle.py @@ -0,0 +1,218 @@ +"""FastSentenceTransformer constructor-redirect lifecycle: +- AutoModel/AutoProcessor/AutoTokenizer.from_pretrained are restored even + when the Transformer constructor raises (try/finally invariant). +- The closure that decides whether to substitute the pre-loaded objects + (`is_requested_model_name`) handles HF repo IDs, local paths, trailing + slashes, pathlib.Path objects, and missing identifiers correctly. +""" + +from __future__ import annotations + +import os +import pathlib +import sys +import types + + +class _FakeAuto: + def __init__(self, name): + self.name = name + self.from_pretrained = self._original + + def _original(self, *args, **kwargs): + return ("orig", self.name, args, kwargs) + + +class _RecordingTransformerOk: + last_calls = None + + def __init__(self, model_name, **kwargs): + from transformers import AutoModel, AutoProcessor, AutoTokenizer + + type(self).last_calls = { + "model": AutoModel.from_pretrained(model_name), + "processor": AutoProcessor.from_pretrained(model_name), + "tokenizer": AutoTokenizer.from_pretrained(model_name), + } + + +class _RaisingTransformer: + def __init__(self, *a, **kw): + from transformers import AutoModel + + AutoModel.from_pretrained(a[0] if a else kw.get("model_name_or_path")) + raise RuntimeError("simulated init failure") + + +def _build_driver(transformer_class): + transformers_mod = types.ModuleType("transformers") + transformers_mod.AutoModel = _FakeAuto("AutoModel") + transformers_mod.AutoProcessor = _FakeAuto("AutoProcessor") + transformers_mod.AutoTokenizer = _FakeAuto("AutoTokenizer") + sys.modules["transformers"] = transformers_mod + + st_root = types.ModuleType("sentence_transformers") + st_models = types.ModuleType("sentence_transformers.models") + st_models.Transformer = transformer_class + sys.modules["sentence_transformers"] = st_root + sys.modules["sentence_transformers.models"] = st_models + + captured = {"calls": None} + + def driver(model_name, model, tokenizer): + from transformers import AutoModel, AutoProcessor, AutoTokenizer + from sentence_transformers.models import Transformer + + def is_requested_model_name(args, kwargs): + requested = None + if args: + requested = args[0] + else: + requested = kwargs.get("pretrained_model_name_or_path") + if requested is None: + requested = kwargs.get("model_name_or_path") + if requested is None: + return False + try: + requested = os.fspath(requested) + expected = os.fspath(model_name) + except (TypeError, ValueError): + return False + if requested == expected: + return True + try: + if os.path.exists(requested) or os.path.exists(expected): + return os.path.abspath(requested) == os.path.abspath(expected) + except (OSError, TypeError, ValueError): + pass + return False + + original_model = AutoModel.from_pretrained + original_processor = AutoProcessor.from_pretrained + original_tokenizer = AutoTokenizer.from_pretrained + + def return_existing_model(*a, **kw): + return model if is_requested_model_name(a, kw) else original_model(*a, **kw) + + def return_existing_tokenizer(*a, **kw): + return ( + tokenizer + if is_requested_model_name(a, kw) + else original_tokenizer(*a, **kw) + ) + + def return_existing_processor(*a, **kw): + return ( + tokenizer + if is_requested_model_name(a, kw) + else original_processor(*a, **kw) + ) + + try: + AutoModel.from_pretrained = return_existing_model + AutoProcessor.from_pretrained = return_existing_processor + AutoTokenizer.from_pretrained = return_existing_tokenizer + t = Transformer(model_name) + captured["calls"] = getattr(type(t), "last_calls", None) + return t + finally: + AutoModel.from_pretrained = original_model + AutoProcessor.from_pretrained = original_processor + AutoTokenizer.from_pretrained = original_tokenizer + + return driver, transformers_mod, captured + + +def test_redirect_substitutes_preloaded_objects_on_match(): + driver, _mod, captured = _build_driver(_RecordingTransformerOk) + sentinel_model = object() + sentinel_tok = object() + driver("sentence-transformers/all-MiniLM-L6-v2", sentinel_model, sentinel_tok) + calls = captured["calls"] + assert calls["model"] is sentinel_model + assert calls["processor"] is sentinel_tok + assert calls["tokenizer"] is sentinel_tok + + +def test_redirect_restored_on_constructor_exception(): + driver, transformers_mod, _ = _build_driver(_RaisingTransformer) + pre_model = transformers_mod.AutoModel.from_pretrained + pre_processor = transformers_mod.AutoProcessor.from_pretrained + pre_tokenizer = transformers_mod.AutoTokenizer.from_pretrained + + try: + driver("model-id", object(), object()) + except RuntimeError: + pass + + assert transformers_mod.AutoModel.from_pretrained is pre_model + assert transformers_mod.AutoProcessor.from_pretrained is pre_processor + assert transformers_mod.AutoTokenizer.from_pretrained is pre_tokenizer + + +def test_redirect_passes_through_for_other_model_names(): + class _OtherNameTransformer: + captured = None + + def __init__(self, model_name, **kw): + from transformers import AutoModel + + type(self).captured = AutoModel.from_pretrained("some-other/aux-model") + + driver, *_ = _build_driver(_OtherNameTransformer) + sentinel = object() + driver("primary/model-id", sentinel, object()) + assert _OtherNameTransformer.captured is not sentinel + assert isinstance(_OtherNameTransformer.captured, tuple) + assert _OtherNameTransformer.captured[0] == "orig" + + +def test_is_requested_model_name_handles_pathlib_path(tmp_path): + target = tmp_path / "model_dir" + target.mkdir() + + class _PathTransformer: + last_calls = None + + def __init__(self, model_name, **kw): + from transformers import AutoModel + + type(self).last_calls = AutoModel.from_pretrained(pathlib.Path(model_name)) + + driver, *_ = _build_driver(_PathTransformer) + sentinel_model = object() + driver(str(target), sentinel_model, object()) + assert _PathTransformer.last_calls is sentinel_model + + +def test_is_requested_model_name_trailing_slash_local_path(tmp_path): + target = tmp_path / "model_dir" + target.mkdir() + + class _SlashTransformer: + last_calls = None + + def __init__(self, model_name, **kw): + from transformers import AutoModel + + type(self).last_calls = AutoModel.from_pretrained(str(target) + "/") + + driver, *_ = _build_driver(_SlashTransformer) + sentinel_model = object() + driver(str(target), sentinel_model, object()) + assert _SlashTransformer.last_calls is sentinel_model + + +def test_is_requested_model_name_returns_false_when_no_identifier(): + captured = {"args": None} + + class _NoNameTransformer: + def __init__(self, model_name, **kw): + from transformers import AutoModel + + captured["args"] = AutoModel.from_pretrained(some_other_kwarg = "x") + + driver, *_ = _build_driver(_NoNameTransformer) + driver("primary/model-id", object(), object()) + assert isinstance(captured["args"], tuple) + assert captured["args"][0] == "orig" diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 541875a3fc..52ff268314 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -23,6 +23,7 @@ from ._utils import ( import inspect import json import os +import threading import types from huggingface_hub import hf_hub_download from typing import Optional @@ -42,6 +43,9 @@ import contextlib import shutil +_CREATE_TRANSFORMER_MODULE_LOCK = threading.RLock() + + def _save_pretrained_torchao( self, save_directory, @@ -1012,37 +1016,128 @@ class FastSentenceTransformer(FastModel): from sentence_transformers.models import Transformer # prevents sentence-transformers from loading the model a second time, thanks Etherl - original_from_pretrained = AutoModel.from_pretrained + # Also redirect AutoProcessor / AutoTokenizer so the Transformer.__init__ + # picks up our pre-fixed tokenizer. On sentence-transformers >=5.4 the + # `tokenizer` attribute became a read-only @property backed by `self.processor`, + # so a post-init assignment raises AttributeError; redirecting the constructor's + # AutoProcessor.from_pretrained call sets self.processor correctly and keeps + # downstream state (input_formatter) consistent. + from transformers import AutoProcessor, AutoTokenizer - def return_existing_model(*args, **kwargs): - return model + def is_requested_model_name(args, kwargs): + requested = None + if args: + requested = args[0] + else: + requested = kwargs.get("pretrained_model_name_or_path") + if requested is None: + requested = kwargs.get("model_name_or_path") + if requested is None: + return False - try: - # Temporarily redirect AutoModel loading to return our pre-loaded model - AutoModel.from_pretrained = return_existing_model + try: + requested = os.fspath(requested) + expected = os.fspath(model_name) + except (TypeError, ValueError) as exception: + logging.debug( + "Unsloth: Could not normalize SentenceTransformer model path: %s", + exception, + ) + return False + if requested == expected: + return True - # Initialize 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}, - ) - finally: - # Restore original functionality immediately - AutoModel.from_pretrained = original_from_pretrained + try: + if os.path.exists(requested) or os.path.exists(expected): + return os.path.abspath(requested) == os.path.abspath(expected) + except (OSError, TypeError, ValueError) as exception: + logging.debug( + "Unsloth: Could not compare SentenceTransformer model paths: %s", + exception, + ) + return False - transformer_module.tokenizer = tokenizer + with _CREATE_TRANSFORMER_MODULE_LOCK: + original_model_from_pretrained = AutoModel.from_pretrained + original_processor_from_pretrained = AutoProcessor.from_pretrained + original_tokenizer_from_pretrained = AutoTokenizer.from_pretrained + + def return_existing_model(*args, **kwargs): + if is_requested_model_name(args, kwargs): + return model + return original_model_from_pretrained(*args, **kwargs) + + def return_existing_tokenizer(*args, **kwargs): + if is_requested_model_name(args, kwargs): + return tokenizer + return original_tokenizer_from_pretrained(*args, **kwargs) + + def return_existing_processor(*args, **kwargs): + if is_requested_model_name(args, kwargs): + return tokenizer + return original_processor_from_pretrained(*args, **kwargs) + + try: + # Temporarily redirect Auto* loading to return our pre-loaded objects + AutoModel.from_pretrained = return_existing_model + AutoProcessor.from_pretrained = return_existing_processor + AutoTokenizer.from_pretrained = return_existing_tokenizer + + transformer_init_params = inspect.signature( + Transformer.__init__ + ).parameters + trust_remote_code_kwargs = {"trust_remote_code": trust_remote_code} + do_lower_case = getattr(tokenizer, "do_lower_case", False) + transformer_kwargs = {"max_seq_length": max_seq_length} + if "do_lower_case" in transformer_init_params: + transformer_kwargs["do_lower_case"] = do_lower_case + if "model_kwargs" in transformer_init_params: + transformer_kwargs["model_kwargs"] = trust_remote_code_kwargs.copy() + transformer_kwargs["config_kwargs"] = ( + trust_remote_code_kwargs.copy() + ) + else: + transformer_kwargs["model_args"] = trust_remote_code_kwargs.copy() + transformer_kwargs["config_args"] = trust_remote_code_kwargs.copy() + if "processor_kwargs" in transformer_init_params: + transformer_kwargs["processor_kwargs"] = ( + trust_remote_code_kwargs.copy() + ) + elif "tokenizer_args" in transformer_init_params: + transformer_kwargs["tokenizer_args"] = ( + trust_remote_code_kwargs.copy() + ) + + # Initialize Transformer + transformer_module = Transformer(model_name, **transformer_kwargs) + finally: + # Restore original functionality immediately + AutoModel.from_pretrained = original_model_from_pretrained + AutoProcessor.from_pretrained = original_processor_from_pretrained + AutoTokenizer.from_pretrained = original_tokenizer_from_pretrained + + # On sentence-transformers >=5.4 `tokenizer` is a read-only property backed + # by `self.processor` (already wired via the redirect above). On older + # versions it's a regular attribute and the explicit assignment is required. + if not isinstance( + getattr(type(transformer_module), "tokenizer", None), property + ): + transformer_module.tokenizer = tokenizer transformer_module.do_lower_case = getattr(tokenizer, "do_lower_case", False) # sentence-transformers only passes along known keys to model.forward + preinit_model_forward_params = getattr( + transformer_module, "model_forward_params", set() + ) 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", + "return_dict", } + transformer_module.model_forward_params |= preinit_model_forward_params # determine max_seq_length if not provided if max_seq_length is None: @@ -1056,7 +1151,11 @@ class FastSentenceTransformer(FastModel): max_seq_length = 512 transformer_module.max_seq_length = max_seq_length - transformer_module.config_keys = ["max_seq_length", "do_lower_case"] + config_keys = list(getattr(transformer_module, "config_keys", []) or []) + for config_key in ("max_seq_length", "do_lower_case"): + if config_key not in config_keys: + config_keys.append(config_key) + transformer_module.config_keys = config_keys transformer_module.save_in_root = True if hasattr(model, "config"): @@ -1064,6 +1163,29 @@ class FastSentenceTransformer(FastModel): return transformer_module + @staticmethod + def _is_transformer_module_ref(class_ref): + if class_ref in { + "sentence_transformers.models.Transformer", + "sentence_transformers.models.transformer.Transformer", + "sentence_transformers.base.modules.transformer.Transformer", + }: + return True + + try: + from sentence_transformers.models import Transformer + from sentence_transformers.util import import_from_string + + module_class = import_from_string(class_ref) + return module_class is Transformer + except (ImportError, AttributeError, TypeError, ValueError) as exception: + logging.debug( + "Unsloth: Could not resolve SentenceTransformer module ref %r: %s", + class_ref, + exception, + ) + return False + @staticmethod def _load_modules( model_name, @@ -1096,7 +1218,7 @@ class FastSentenceTransformer(FastModel): "name", str(module_config.get("idx", len(modules))) ) - if class_ref == "sentence_transformers.models.Transformer": + if FastSentenceTransformer._is_transformer_module_ref(class_ref): transformer_module = ( FastSentenceTransformer._create_transformer_module( model_name, @@ -1299,7 +1421,10 @@ class FastSentenceTransformer(FastModel): if hasattr(model, "__getitem__"): inner_model = model[0].auto_model compiled = torch.compile(inner_model, mode = mode) - model[0].auto_model = compiled + if isinstance(getattr(type(model[0]), "auto_model", None), property): + model[0].model = compiled + else: + model[0].auto_model = compiled # Fix for accelerate unwrap_model bug: # When SentenceTransformer contains a compiled inner model, # accelerate checks has_compiled_regions() which returns True, @@ -1933,8 +2058,15 @@ class FastSentenceTransformer(FastModel): "Unsloth: Re-enabling torch.compile since gradient checkpointing is not supported" ) - # Re-assign the peft model back to the transformer module - transformer_module.auto_model = peft_model + # Re-assign the peft model back to the transformer module. + # On sentence-transformers >=5.4 `auto_model` is a read-only property + # backed by `self.model`, so write to the backing attribute there. + if isinstance( + getattr(type(transformer_module), "auto_model", None), property + ): + transformer_module.model = peft_model + else: + transformer_module.auto_model = peft_model # Store compile info for auto-compile at trainer time # torch.compile is deferred until training starts so we can check max_steps @@ -1980,8 +2112,15 @@ class FastSentenceTransformer(FastModel): **kwargs, ) - # re-assign the peft model back to the transformer module - transformer_module.auto_model = peft_model + # re-assign the peft model back to the transformer module. + # On sentence-transformers >=5.4 `auto_model` is a read-only property + # backed by `self.model`, so write to the backing attribute there. + if isinstance( + getattr(type(transformer_module), "auto_model", None), property + ): + transformer_module.model = peft_model + else: + transformer_module.auto_model = peft_model return model else: return FastModel.get_peft_model(