From 07bafe2c9336c932ab95324aa36cd9b374e89ed7 Mon Sep 17 00:00:00 2001 From: electroglyph Date: Thu, 1 Jan 2026 03:50:36 -0800 Subject: [PATCH] fix mpnet gradient checkpointing for torch >= 2.9 --- unsloth/models/sentence_transformer.py | 15 +++++---------- 1 file changed, 5 insertions(+), 10 deletions(-) diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 169197865d..65d6d8075a 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 = False, + 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,14 +716,9 @@ 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]