Update torchao save (#3679)

* Update torchao save

* up

* up

* up

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Apply suggestion from @danielhanchen

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
This commit is contained in:
Scott Roy 2025-12-12 04:07:02 -08:00 committed by GitHub
commit c91e99370b

View file

@ -42,6 +42,7 @@ import re
from transformers.models.llama.modeling_llama import logger
from .tokenizer_utils import fix_sentencepiece_gguf
from .models.loader_utils import get_model_name
from .models._utils import _convert_torchao_model
from .ollama_template_mappers import OLLAMA_TEMPLATES, MODEL_TO_OLLAMA_TEMPLATE_MAPPER
from transformers import ProcessorMixin
from huggingface_hub import HfApi
@ -2734,11 +2735,35 @@ def unsloth_generic_push_to_hub_merged(
gc.collect()
def unsloth_save_pretrained_torchao(
self,
def _unsloth_save_torchao_with_attached_config(
model,
save_directory: Union[str, os.PathLike],
tokenizer = None,
torchao_config = None,
tokenizer,
push_to_hub: bool = False,
token: Optional[Union[str, bool]] = None,
):
"""Save a QAT-trained model by converting fake-quantized weights to real quantized weights."""
# Convert QAT fake-quantized weights to real quantized weights
_convert_torchao_model(model)
# TorchAO does not support safe_serialization reliably
safe_serialization = False
if push_to_hub:
model.push_to_hub(
save_directory, safe_serialization = safe_serialization, token = token
)
tokenizer.push_to_hub(save_directory, token = token)
else:
model.save_pretrained(save_directory, safe_serialization = safe_serialization)
tokenizer.save_pretrained(save_directory)
def _unsloth_save_torchao_with_given_config(
model,
save_directory: Union[str, os.PathLike],
tokenizer,
torchao_config,
push_to_hub: bool = False,
token: Optional[Union[str, bool]] = None,
):
@ -2749,23 +2774,26 @@ def unsloth_save_pretrained_torchao(
`torchao_config` (TorchAOBaseConfig): configuration for torchao quantization, full list: https://docs.pytorch.org/ao/main/api_ref_quantization.html#inference-apis-for-quantize
`push_to_hub` (bool): whether to push the checkpoint to huggingface hub or save locally
"""
if push_to_hub:
assert token is not None, "Unsloth: Please specify a token for uploading!"
assert (
torchao_config is not None
), "Unsloth: Please specify a torchao_config for post-training quantization!"
# first merge the lora weights
arguments = dict(locals())
arguments["model"] = self
arguments["tokenizer"] = tokenizer
arguments["push_to_hub"] = False # We save ourselves
arguments["save_method"] = "merged_16bit" # Must be 16bit
del arguments["self"]
del arguments["torchao_config"]
if token is None and push_to_hub:
token = get_token()
if not isinstance(self, PeftModelForCausalLM) and not isinstance(self, PeftModel):
self.save_pretrained(save_directory)
if not isinstance(model, PeftModelForCausalLM) and not isinstance(model, PeftModel):
model.save_pretrained(save_directory)
tokenizer.save_pretrained(save_directory)
else:
unsloth_generic_save(**arguments)
for _ in range(3):
gc.collect()
@ -2778,26 +2806,20 @@ def unsloth_save_pretrained_torchao(
)
from torchao import quantize_
if torchao_config is None:
from torchao.quantization import Int8DynamicActivationInt8WeightConfig
print(
"Unsloth: You did not specify a `torchao_config`, so defaulting to `Int8DynamicActivationInt8WeightConfig`"
)
torchao_config = Int8DynamicActivationInt8WeightConfig()
quantization_config = TorchAoConfig(quant_type = torchao_config)
# Determine if this is a VLM
is_vlm = False
if hasattr(self, "config") and hasattr(self.config, "architectures"):
if hasattr(model, "config") and hasattr(model.config, "architectures"):
is_vlm = any(
x.endswith(("ForConditionalGeneration", "ForVisionText2Text"))
for x in self.config.architectures
for x in model.config.architectures
)
is_vlm = is_vlm or hasattr(self.config, "vision_config")
is_vlm = is_vlm or hasattr(model.config, "vision_config")
auto_model = AutoModelForImageTextToText if is_vlm else AutoModelForCausalLM
auto_processor = AutoProcessor if is_vlm else AutoTokenizer
tokenizer = auto_processor.from_pretrained(arguments["save_directory"])
tokenizer = auto_processor.from_pretrained(save_directory)
# TorchAO must only use bfloat16 for loading (float16 fails)
if HAS_TORCH_DTYPE:
@ -2805,8 +2827,9 @@ def unsloth_save_pretrained_torchao(
else:
kwargs = {"dtype": torch.bfloat16}
model = auto_model.from_pretrained(
arguments["save_directory"],
# Reload with quantization applied
quantized_model = auto_model.from_pretrained(
save_directory,
device_map = "auto",
quantization_config = quantization_config,
**kwargs,
@ -2817,25 +2840,92 @@ def unsloth_save_pretrained_torchao(
# TorchAO does not support safe_serialization right now 0.14.0 seems broken!
safe_serialization = Version(importlib_version("torchao")) > Version("0.14.0")
safe_serialization = False
if push_to_hub:
if token is None and push_to_hub:
token = get_token()
model.push_to_hub(
quantized_model.push_to_hub(
torchao_save_directory, safe_serialization = safe_serialization, token = token
)
tokenizer.push_to_hub(torchao_save_directory, token = token)
else:
model.save_pretrained(
quantized_model.save_pretrained(
torchao_save_directory, safe_serialization = safe_serialization
)
tokenizer.save_pretrained(torchao_save_directory)
# Clean up the intermediate unquantized model
if os.path.exists(save_directory):
try:
import shutil
shutil.rmtree(save_directory)
except:
pass
def unsloth_save_pretrained_torchao(
self,
save_directory: Union[str, os.PathLike],
tokenizer = None,
torchao_config = None,
push_to_hub: bool = False,
token: Optional[Union[str, bool]] = None,
):
"""Saves a torchao quantized model checkpoint.
This function handles two mutually exclusive workflows:
1. **QAT (Quantization-Aware Training)**: If the model was trained with `qat_scheme`
parameter, do NOT pass `torchao_config`. The function will convert the QAT
fake-quantized weights to real quantized weights and save directly.
2. **PTQ (Post-Training Quantization)**: If you want to apply quantization to a
regular model, pass a `torchao_config`. The model must NOT have been trained
with `qat_scheme`.
Args:
`save_directory`: local folder path or huggingface hub ID when `push_to_hub` is True
`tokenizer`: the tokenizer to save alongside the model
`torchao_config` (TorchAOBaseConfig): configuration for torchao quantization.
Required for PTQ, must be None for QAT models.
Options: https://docs.pytorch.org/ao/main/api_ref_quantization.html#inference-apis-for-quantize
`push_to_hub` (bool): whether to push to huggingface hub or save locally
`token`: HuggingFace token for pushing to hub
"""
if token is None and push_to_hub:
token = get_token()
has_qat_config = (
hasattr(self, "_torchao_config") and self._torchao_config is not None
)
if torchao_config is not None:
# PTQ path: user provided a config, model must NOT have QAT config
assert not has_qat_config, (
"Unsloth: You passed `torchao_config` but this model was trained with `qat_scheme`. "
"For QAT models, do not pass `torchao_config` - the quantization config is already "
"attached to the model from training."
)
_unsloth_save_torchao_with_given_config(
model = self,
save_directory = save_directory,
tokenizer = tokenizer,
torchao_config = torchao_config,
push_to_hub = push_to_hub,
token = token,
)
else:
# QAT path: no config provided, model must have QAT config
assert has_qat_config, (
"Unsloth: No `torchao_config` provided and model was not trained with `qat_scheme`. "
"Either train with `qat_scheme` parameter, or provide a `torchao_config` for "
"post-training quantization."
)
_unsloth_save_torchao_with_attached_config(
model = self,
save_directory = save_directory,
tokenizer = tokenizer,
push_to_hub = push_to_hub,
token = token,
)
for _ in range(3):
gc.collect()