Update save.py
This commit is contained in:
parent
34998914ac
commit
44659a05a3
1 changed files with 23 additions and 7 deletions
|
|
@ -213,7 +213,7 @@ def unsloth_save_model(
|
|||
pass
|
||||
save_pretrained_settings["tags"] = tags
|
||||
|
||||
if (save_method == "lora") and push_to_hub:
|
||||
if ((save_method == "lora") or (save_method == "merged_4bit")) and push_to_hub:
|
||||
if token is None:
|
||||
raise RuntimeError(
|
||||
"Unsloth: Pushing to HF requires a token. Pass `token = 'hf_....'`\n"\
|
||||
|
|
@ -249,17 +249,14 @@ def unsloth_save_model(
|
|||
tags = tags,
|
||||
)
|
||||
pass
|
||||
return save_directory
|
||||
pass
|
||||
|
||||
# Update model tag
|
||||
username = ""
|
||||
if push_to_hub:
|
||||
username = upload_to_huggingface(
|
||||
# Update model tag
|
||||
_ = upload_to_huggingface(
|
||||
model, save_directory, token,
|
||||
"finetuned", "trl", file_location = None,
|
||||
old_username = None,
|
||||
)
|
||||
return save_directory
|
||||
pass
|
||||
|
||||
# Tokenizer has different saving arguments
|
||||
|
|
@ -310,6 +307,16 @@ def unsloth_save_model(
|
|||
if save_method != "lora": print(" This might take 10 minutes for Llama-7b...", end = "")
|
||||
|
||||
model.save_pretrained(**save_pretrained_settings)
|
||||
|
||||
# Update model tag
|
||||
if push_to_hub:
|
||||
_ = upload_to_huggingface(
|
||||
model, save_pretrained_settings["save_directory"], token,
|
||||
"finetuned", "trl", file_location = None,
|
||||
old_username = None,
|
||||
)
|
||||
pass
|
||||
|
||||
print(" Done.")
|
||||
return save_directory
|
||||
pass
|
||||
|
|
@ -474,6 +481,15 @@ def unsloth_save_model(
|
|||
model.config = old_config
|
||||
print("Done.")
|
||||
|
||||
# Update model tag
|
||||
if push_to_hub:
|
||||
_ = upload_to_huggingface(
|
||||
model, save_pretrained_settings["save_directory"], token,
|
||||
"finetuned", "trl", file_location = None,
|
||||
old_username = username,
|
||||
)
|
||||
pass
|
||||
|
||||
# Print location
|
||||
if push_to_hub:
|
||||
print(f"Saved to https://huggingface.co/{username}/{save_directory.lstrip('/')}")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue