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:
parent
a458d59e91
commit
b701d743ee
1 changed files with 27 additions and 3 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue