From 3acea66cc54a2c2a7b1d710c502d342fb8a8d2b2 Mon Sep 17 00:00:00 2001 From: electroglyph Date: Tue, 16 Dec 2025 03:05:05 -0800 Subject: [PATCH] add save_pretrained_merged method which gets modules and config --- unsloth/models/sentence_transformer.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index 3e8850cf39..0be988483e 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -17,6 +17,7 @@ import torch import inspect import json import os +import types from huggingface_hub import hf_hub_download @@ -219,6 +220,23 @@ class FastSentenceTransformer(FastModel): normalize_module = Normalize() modules = [transformer_module, pooling_module, normalize_module] st_model = SentenceTransformer(modules = modules) + + def _save_pretrained_merged(self, save_directory, **kwargs): + # sentence-transformers config and modules only get saved if we call save_pretrained + self.save_pretrained(save_directory) + + # Remove LoRA adapters since we are saving the merged model + for file in ["adapter_model.safetensors", "adapter_config.json"]: + try: + os.remove(os.path.join(save_directory, file)) + except: + pass + + # save merged weights + tokenizer = kwargs.pop("tokenizer", self.tokenizer) + self[0].auto_model.save_pretrained_merged(save_directory, tokenizer=tokenizer, **kwargs) + + st_model.save_pretrained_merged = types.MethodType(_save_pretrained_merged, st_model) return st_model @staticmethod