Fix/save torchao model loading logic (#3621)

* make loading gpt-oss-BF16 faster. Linked to unsloth-zoo PR #314

* fix model loading and clean merged model directory

* revert default quant

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

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

* revert mapper.py

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Roland Tannous 2025-11-20 10:02:31 +02:00 committed by GitHub
commit b701d743ee

View file

@ -2769,7 +2769,13 @@ def unsloth_save_pretrained_torchao(
for _ in range(3):
gc.collect()
from transformers import AutoModel, AutoTokenizer, TorchAoConfig
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
TorchAoConfig,
AutoModelForImageTextToText,
AutoProcessor,
)
from torchao import quantize_
if torchao_config is None:
@ -2781,14 +2787,25 @@ def unsloth_save_pretrained_torchao(
torchao_config = Int8DynamicActivationInt8WeightConfig()
quantization_config = TorchAoConfig(quant_type = torchao_config)
tokenizer = AutoTokenizer.from_pretrained(arguments["save_directory"])
is_vlm = False
if hasattr(self, "config") and hasattr(self.config, "architectures"):
is_vlm = any(
x.endswith(("ForConditionalGeneration", "ForVisionText2Text"))
for x in self.config.architectures
)
is_vlm = is_vlm or hasattr(self.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"])
# TorchAO must only use bfloat16 for loading (float16 fails)
if HAS_TORCH_DTYPE:
kwargs = {"torch_dtype": torch.bfloat16}
else:
kwargs = {"dtype": torch.bfloat16}
model = AutoModel.from_pretrained(
model = auto_model.from_pretrained(
arguments["save_directory"],
device_map = "auto",
quantization_config = quantization_config,
@ -2812,6 +2829,13 @@ def unsloth_save_pretrained_torchao(
torchao_save_directory, safe_serialization = safe_serialization
)
tokenizer.save_pretrained(torchao_save_directory)
if os.path.exists(save_directory):
try:
import shutil
shutil.rmtree(save_directory)
except:
pass
for _ in range(3):
gc.collect()