From 680d43a488362f9efa380d27176a220ad25d3cc2 Mon Sep 17 00:00:00 2001 From: Etherll <61019402+Etherll@users.noreply.github.com> Date: Tue, 5 May 2026 14:15:54 +0300 Subject: [PATCH] Fix FastSentenceTransformer loading with newer sentence-transformers (#5259) * Fix FastSentenceTransformer compatibility with sentence-transformers 5.4 * Support varied Transformer init signatures Detect Transformer.__init__ parameters and build init kwargs accordingly so trust_remote_code and other args are passed using the correct names. Instead of unconditionally using model_args/config_args, the code now inspects the constructor to decide between model_kwargs/config_kwargs vs model_args/config_args and also sets processor_kwargs or tokenizer_args when present. Initializes Transformer with constructed transformer_kwargs (including max_seq_length) to improve compatibility with different Transformer implementations. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Harden SentenceTransformer path and module checks * Scrub .github/workflows for staging push (matches staging base) * Guard auto_model write in FastSentenceTransformer._apply_torch_compile On sentence-transformers >=5.4 Transformer.auto_model is a read-only @property backed by self.model, so a direct assignment raises AttributeError. The two get_peft_model paths already guard the write with isinstance(getattr(type(...), "auto_model", None), property); the auto-compile path missed the same guard, which broke the default trainer path whenever max_steps >= _compile_threshold. * Add tests for FastSentenceTransformer property guards * Tighten FastSentenceTransformer redirect lifecycle tests Drop a duplicate assertion-less case, remove dead AST extraction helper, and trim unused imports. The remaining six tests cover substitution on match, restoration on constructor exception, passthrough for unrelated names, pathlib.Path normalisation, trailing slash handling, and the no-identifier guard. * Sync .github/workflows with upstream author branch * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Avoid sharing trust_remote_code kwargs dict across constructor buckets In FastSentenceTransformer._create_transformer_module, the same trust_remote_code_kwargs dict was being assigned to model_kwargs, config_kwargs, and processor_kwargs (or model_args / config_args / tokenizer_args) on the Transformer constructor. transformers' from_pretrained code paths (configuration_utils, auto_factory, processing_auto, etc.) call kwargs.pop("trust_remote_code", ...) on the dict they receive, which would drain the shared object and silently strip trust_remote_code from the other buckets. Pass an independent copy to each bucket so subsequent buckets and any pass-through auxiliary loads still see trust_remote_code. * Wire do_lower_case and return_dict through Transformer init for ST 5.4 In FastSentenceTransformer._create_transformer_module: - When Transformer.__init__ accepts do_lower_case (ST 5.4+), pass the unsloth tokenizer's do_lower_case as a constructor kwarg. The existing post-init attribute assignment alone is too late: ST 5.4's __init__ uses do_lower_case to install a Lowercase normalizer on tokenizer.backend_tokenizer.normalizer, which is not re-applied if we only set the attribute after construction. The post-init line is preserved untouched for older ST versions. - Add return_dict to the manually completed model_forward_params set so wrapped models with forward(*args, **kwargs) signatures keep ST's forced dict-like output safety net. ST 5.4's own __init__ unions the forward signature with the same set plus return_dict; the previous override silently dropped it. * Preserve flash-attention forward keys when wrapping ST 5.4 Transformer Sentence-transformers 5.4's Transformer.__init__ calls _can_flatten_inputs() during construction, which augments self.model_forward_params with cu_seq_lens_q, cu_seq_lens_k, max_length_q, max_length_k, seq_idx whenever feature-extraction with text modality, the torch backend, flash-attention 2, and varlen flash-attn support are all available. The post-init override of transformer_module.model_forward_params used to replace the attribute outright, silently dropping those keys so ST's preprocess() filter stripped flash-attn kwargs before reaching model.forward. Snapshot the constructor-populated set first, leave the existing overwrite intact for the forward-signature plus tokenizer keys, and union the snapshot back in so flash-attn forwarding keeps working on ST 5.4. For older sentence-transformers releases the attribute is absent and getattr returns an empty set, leaving behavior unchanged. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han --- ...sentence_transformer_redirect_lifecycle.py | 218 ++++++++++++++++++ unsloth/models/sentence_transformer.py | 187 +++++++++++++-- 2 files changed, 381 insertions(+), 24 deletions(-) create mode 100644 tests/python/test_fast_sentence_transformer_redirect_lifecycle.py 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(