-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsave_model.py
More file actions
45 lines (35 loc) · 1.34 KB
/
Copy pathsave_model.py
File metadata and controls
45 lines (35 loc) · 1.34 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
import argparse
from transformers import AutoModelForCausalLM, PreTrainedModel, PreTrainedTokenizer, AutoTokenizer
def push_model(model: PreTrainedModel, tokenizer: PreTrainedTokenizer, path: str) -> None:
model.push_to_hub(path)
tokenizer.push_to_hub(path)
def load_model(path: str) -> tuple[PreTrainedModel, PreTrainedTokenizer]:
model = AutoModelForCausalLM.from_pretrained(
path,
)
tokenizer = AutoTokenizer.from_pretrained(path)
return model, tokenizer
def main():
parser = argparse.ArgumentParser(description="Save and Load Model Script")
parser.add_argument(
"--model_path", type=str, required=True, help="Path to the model directory"
)
parser.add_argument(
"--remote_path", type=str, required=True, help="Remote path to push the model"
)
args = parser.parse_args()
model, tokenizer = load_model(args.model_path)
print(f"Model loaded from {args.model_path}")
print(
"Do you want to proceed with uploading the pruned model to the Hugging Face Hub? (y/n): ",
end="",
)
choice = input().strip().lower()
if choice == "y":
print("Uploading model to Hugging Face Hub...")
push_model(model, tokenizer, args.remote_path)
print("Upload complete.")
else:
print("Upload skipped.")
if __name__ == "__main__":
main()