From b2742f2543cf9d0304e65ae5525997ac4fbaee41 Mon Sep 17 00:00:00 2001 From: electroglyph Date: Wed, 7 Jan 2026 19:17:16 -0800 Subject: [PATCH] add save_pretrained_torchao --- unsloth/models/sentence_transformer.py | 108 +++++++++++++++++++++++++ 1 file changed, 108 insertions(+) diff --git a/unsloth/models/sentence_transformer.py b/unsloth/models/sentence_transformer.py index e57ada7807..d3617f630a 100644 --- a/unsloth/models/sentence_transformer.py +++ b/unsloth/models/sentence_transformer.py @@ -34,6 +34,106 @@ from transformers import AutoModel, AutoConfig from transformers.models.auto.auto_factory import _get_model_class import tempfile from huggingface_hub import HfApi, get_token +from ..save import unsloth_save_pretrained_torchao +import contextlib +import shutil + + +def _save_pretrained_torchao( + self, + save_directory, + tokenizer=None, + torchao_config=None, + push_to_hub=False, + token=None, +): + self.save_pretrained(save_directory) + + # grab inner model + inner_model = self[0].auto_model + if hasattr(inner_model, "_orig_mod"): + inner_model = inner_model._orig_mod + + # merge LoRA first + if hasattr(inner_model, "merge_and_unload"): + inner_model = inner_model.merge_and_unload() + + # confirm Transformer path + transformer_path = "0_Transformer" + modules_path = os.path.join(save_directory, "modules.json") + if os.path.exists(modules_path): + try: + with open(modules_path, "r") as f: + modules = json.load(f) + for m in modules: + if m.get("type", "").endswith("Transformer"): + transformer_path = m.get("path", "") + break + except: + pass + + transformer_dir = os.path.join(save_directory, transformer_path) + transformer_dir = os.path.abspath(transformer_dir) + + if tokenizer is None: + tokenizer = self.tokenizer + + @contextlib.contextmanager + def patch_unsloth_save(): + original_causal = transformers.AutoModelForCausalLM + original_rmtree = shutil.rmtree + # unsloth_save_pretrained_torchao expects AutoModelForCausalLM + transformers.AutoModelForCausalLM = transformers.AutoModel + # prevent unsloth from deleting the unquantized model directory + shutil.rmtree = lambda *args, **kwargs: None + try: + yield + finally: + # unpatch + transformers.AutoModelForCausalLM = original_causal + shutil.rmtree = original_rmtree + + with patch_unsloth_save(): + unsloth_save_pretrained_torchao( + inner_model, + transformer_dir, + tokenizer=tokenizer, + torchao_config=torchao_config, + push_to_hub=push_to_hub, + token=token, + ) + + # avoid `0_Transformer-torchao`, it was either this or fix modules.json + torchao_dir = transformer_dir + "-torchao" + if os.path.exists(torchao_dir): + if not os.path.exists(transformer_dir): + os.makedirs(transformer_dir, exist_ok=True) + + # move contents + for item in os.listdir(torchao_dir): + s = os.path.join(torchao_dir, item) + d = os.path.join(transformer_dir, item) + if os.path.isdir(s): + shutil.copytree(s, d, dirs_exist_ok=True) + else: + shutil.copy2(s, d) + + # remove torchao dir + shutil.rmtree(torchao_dir) + + # remove conflicting safetensors if we brought in bin + if os.path.exists(os.path.join(transformer_dir, "pytorch_model.bin")): + safetensors_path = os.path.join(transformer_dir, "model.safetensors") + if os.path.exists(safetensors_path): + try: + os.remove(safetensors_path) + except: + pass + + try: + FastSentenceTransformer._add_unsloth_branding(save_directory) + except: + pass class FastSentenceTransformer(FastModel): @@ -881,6 +981,10 @@ class FastSentenceTransformer(FastModel): _save_pretrained_merged, st_model ) + st_model.save_pretrained_torchao = types.MethodType( + _save_pretrained_torchao, st_model + ) + def _push_to_hub_merged(self, repo_id, **push_kwargs): hub_token = push_kwargs.get("token", None) or get_token() if hub_token is None: @@ -1075,6 +1179,10 @@ class FastSentenceTransformer(FastModel): _save_pretrained_merged, st_model ) + st_model.save_pretrained_torchao = types.MethodType( + _save_pretrained_torchao, st_model + ) + def _push_to_hub_merged(self, repo_id, **kwargs): token = kwargs.get("token", None) or get_token() if token is None: