From 0af3a1b14357f1cd4712eaf00122f40e3d2d89c7 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 18 Feb 2024 19:44:59 +1100 Subject: [PATCH] Update save.py --- unsloth/save.py | 51 +++++++++++++++++++++++++------------------------ 1 file changed, 26 insertions(+), 25 deletions(-) diff --git a/unsloth/save.py b/unsloth/save.py index 6ab6768da1..06d0c6e5ec 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -250,23 +250,6 @@ def unsloth_save_model( ) pass - # If push_to_hub, we must remove the .../ part of a repo - # username = None - # if push_to_hub and "/" in save_directory: - - # # +1 solves absolute path issues - # username = save_directory[:save_directory.find("/")] - # new_save_directory = save_directory[save_directory.find("/")+1:] - - # logger.warning_once( - # f"Unsloth: You are pushing to hub, but you passed your HF username = {username}.\n"\ - # f"We shall truncate {save_directory} to {new_save_directory}" - # ) - - # save_pretrained_settings["save_directory"] = new_save_directory - # save_directory = new_save_directory - # pass - # Tokenizer has different saving arguments tokenizer_save_settings = \ { @@ -316,6 +299,24 @@ def unsloth_save_model( return save_directory pass + # If push_to_hub, we must remove the .../ part of a repo + username = None + if push_to_hub and "/" in save_directory: + + # +1 solves absolute path issues + username = save_directory[:save_directory.find("/")] + new_save_directory = save_directory[save_directory.find("/")+1:] + + logger.warning_once( + f"Unsloth: You are pushing to hub, but you passed your HF username = {username}.\n"\ + f"We shall truncate {save_directory} to {new_save_directory}" + ) + + save_pretrained_settings["save_directory"] = new_save_directory + tokenizer_save_settings ["save_directory"] = new_save_directory + save_directory = new_save_directory + pass + print("Unsloth: Merging 4bit and LoRA weights to 16bit...") # Determine max RAM usage minus sharding @@ -1040,8 +1041,6 @@ def patch_saving_functions(model): import types from typing import Callable, Optional, Union, List - if hasattr(model, "_original_push_to_hub"): return - # First check if this has already been called, and revert it original_model = model while True: @@ -1052,6 +1051,8 @@ def patch_saving_functions(model): if hasattr(original_model, "save_pretrained_merged"): del original_model.save_pretrained_merged if hasattr(original_model, "push_to_hub_gguf"): del original_model.push_to_hub_gguf if hasattr(original_model, "save_pretrained_gguf"): del original_model.save_pretrained_gguf + else: + original_model._original_push_to_hub = original_model.push_to_hub pass if hasattr(original_model, "model"): original_model = original_model.model @@ -1059,12 +1060,11 @@ def patch_saving_functions(model): pass # And now re add our saving methods! - original_push_to_hub = model.push_to_hub + original_push_to_hub = model._original_push_to_hub signature = str(inspect.signature(original_push_to_hub)).replace("NoneType", "None") signature = signature[1:] signature = re.sub("", "torch.save", signature) docs = original_push_to_hub.__doc__.encode("utf-8").decode("utf-8") - model._original_push_to_hub = original_push_to_hub push_to_hub_text = f'''def unsloth_push_to_hub(self, {signature}: """ @@ -1093,12 +1093,13 @@ def patch_saving_functions(model): if not hasattr(original_model, "_original_push_to_hub"): original_model._original_push_to_hub = original_model.push_to_hub - original_model.push_to_hub = types.MethodType(unsloth_push_to_hub, original_model) - - if hasattr(original_model, "add_model_tags"): - original_model.add_model_tags(["unsloth",]) pass + original_model.push_to_hub = types.MethodType(unsloth_push_to_hub, original_model) + if hasattr(original_model, "add_model_tags"): + original_model.add_model_tags(["unsloth",]) + pass + if hasattr(original_model, "model"): original_model = original_model.model else: break pass