add save_pretrained_torchao

This commit is contained in:
electroglyph 2026-01-07 19:17:16 -08:00
commit b2742f2543

View file

@ -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: