diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 65d6d8075a..4bb49f72b0 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -159,7 +159,7 @@ class FastSentenceTransformer(FastModel): attention_mask, head_mask[i] if head_mask is not None else None, position_bias, - use_reentrant = True, # fix for torch 2.9 + use_reentrant = True, # fix for torch 2.9 ) else: # original code from here on @@ -702,12 +702,12 @@ class FastSentenceTransformer(FastModel): isinstance(st_device, str) and st_device in ["auto", "sequential"] ): st_device = None - + # this was added because when loading for inference it was defaulting to float32 # propagate dtype to model_kwargs, default to "auto" model_kwargs = kwargs.get("model_kwargs", {}) model_kwargs["dtype"] = dtype if dtype is not None else "auto" - + # filter kwargs for SentenceTransformer st_kwargs = { "device": st_device, @@ -716,9 +716,14 @@ class FastSentenceTransformer(FastModel): "revision": revision, "model_kwargs": model_kwargs, } - + # add other known kwargs if present - known_keys = ["cache_folder", "truncate_dim", "tokenizer_kwargs", "config_kwargs"] + known_keys = [ + "cache_folder", + "truncate_dim", + "tokenizer_kwargs", + "config_kwargs", + ] for k in known_keys: if k in kwargs: st_kwargs[k] = kwargs[k]