-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
111 lines (107 loc) · 4.22 KB
/
Copy pathmain.py
File metadata and controls
111 lines (107 loc) · 4.22 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
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
"""Main interface to run inference or fine-tuning"""
import os
import argparse
import torch
from src.inference import inference, dataset_inference
from src.training import finetuning
from src.eda import eda
from src.tune import hpo_tune
if __name__ == "__main__":
# Set any environment variables and PyTorch performance settings
if torch.cuda.is_available():
torch.backends.cudnn.benchmark = True
os.environ["OMP_NUM_THREADS"] = f"{os.cpu_count()}"
os.environ["TOKENIZERS_PARALLELISM"] = "false"
torch.manual_seed(1337)
parser = argparse.ArgumentParser()
parser.add_argument(
"--tokenizer", type=str, default="hf-internal-testing/llama-tokenizer"
)
parser.add_argument(
"--model",
type=str,
# default="TinyLlama/TinyLlama-1.1B-intermediate-step-1431k-3T",
default="TinyLlama/TinyLlama-1.1B-intermediate-step-1195k-token-2.5T",
# default="codellama/CodeLlama-7b-hf",
)
parser.add_argument("--temperature", "-t", type=float, default=0.0)
parser.add_argument("--n-tokens-to-generate", "-N", type=int, default=500)
parser.add_argument("--dataset", type=str, default="openai_humaneval")
parser.add_argument("--test-dataset", type=str, default="openai_humaneval")
parser.add_argument(
"--train-dataset",
type=str,
default="iamtarun/python_code_instructions_18k_alpaca",
)
parser.add_argument(
"--prompt",
"-p",
type=str,
default="Question: Write the fibonacci sequence in Python. \nAnswer: def fibonacci(n):",
)
parser.add_argument(
"--instruction-prompt",
"-ip",
type=str,
default="Please finish coding the Python script provided below without performing \
any additional tasks such as testing or writing explanations.\n",
)
parser.add_argument("--inference", action="store_true")
parser.add_argument("--inference-evaluate", action="store_true")
parser.add_argument("--finetuning", action="store_true")
parser.add_argument("--hyperparameter-tune", action="store_true")
parser.add_argument("--eda", action="store_true")
parser.add_argument("--rank", type=int, default=32)
parser.add_argument("--alpha", type=float, default=1.0)
parser.add_argument("--layers", type=int, default=4)
parser.add_argument("--epochs", type=int, default=10)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--dropout", type=float, default=0.1)
parser.add_argument("--lr", type=float, default=1e-5)
parser.add_argument("--lora-checkpoint-path", type=str)
parser.add_argument("--debug", "-d", action="store_true")
args = parser.parse_args()
# Run inference on a single prompt
if args.inference:
# Disable Autograd when running inference
with torch.no_grad():
inference(
args.model,
args.tokenizer,
args.prompt,
lora_checkpoint_path=args.lora_checkpoint_path,
)
# Evaluate inference over the whole dataset
elif args.inference_evaluate:
# Disable Autograd when running inference
with torch.no_grad():
dataset_inference(
args.model,
args.tokenizer,
args.test_dataset,
temperature=args.temperature,
N=args.n_tokens_to_generate,
lora_checkpoint_path=args.lora_checkpoint_path,
instruction_prompt=args.instruction_prompt,
)
elif args.finetuning:
finetuning(
args.model,
args.tokenizer,
args.train_dataset,
epochs=args.epochs,
batch_size=args.batch_size,
rank=args.rank,
alpha=args.alpha,
layers=args.layers,
dropout=args.dropout,
lr=args.lr,
instruction_prompt=args.instruction_prompt,
)
elif args.hyperparameter_tune:
hpo_tune(args.model, args.tokenizer, args.train_dataset)
elif args.eda:
eda(args.model, args.tokenizer, args.train_dataset, training=True)
eda(args.model, args.tokenizer, args.test_dataset, training=False)
else:
raise ValueError("Invalid mode.")