diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 85eea2bf3f..58b6b99107 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -702,10 +702,28 @@ 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, + "trust_remote_code": trust_remote_code, + "token": token, + "revision": revision, + "model_kwargs": model_kwargs, + } + + # add other known kwargs if present + known_keys = ["cache_folder", "truncate_dim", "tokenizer_kwargs", "config_kwargs"] + for k in known_keys: + if k in kwargs: + st_kwargs[k] = kwargs[k] - st_model = SentenceTransformer( - model_name, device = st_device, trust_remote_code = trust_remote_code - ) + st_model = SentenceTransformer(model_name, **st_kwargs) return st_model if "auto_model" not in kwargs: