-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathtrain_arithmetic.py
More file actions
176 lines (139 loc) · 6.12 KB
/
Copy pathtrain_arithmetic.py
File metadata and controls
176 lines (139 loc) · 6.12 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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
import torch
from torch.utils.data import DataLoader
from transformers import (
RobertaTokenizer,
RobertaForSequenceClassification,
AdamW,
get_linear_schedule_with_warmup,
TrainingArguments,
Trainer
)
import time
from datasets import load_dataset
from tqdm.auto import tqdm
import numpy as np
from peft import get_peft_model, LoraConfig, TaskType
import argparse
import warnings
import os
from datetime import datetime
import json
import yaml
import atexit
import wandb
from utils.data_utils import *
from models import *
from utils.misc import *
import os
os.environ['MASTER_ADDR'] = 'localhost'
os.environ['MASTER_PORT'] = '12355'
def create_run_directory(args):
"""Create a directory structure for the current training run."""
# Create base directory for all runs
base_dir = "experiments/arithmetic"
# Create timestamp for unique run identification
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
# Create model name directory (simplified name)
model_name = args.model.split('/')[-1]
# Create run-specific directory with relevant parameters
run_name = f"rank_{args.lora_r}_lr{args.lr}_alpha_{args.lora_alpha}_train_{args.dataset_split.replace('[:','').replace(']','')}"
# Final directory structure: experiments/model_name/training_type/YYYYMMDD_HHMMSS_parameters
run_dir = os.path.join(base_dir, model_name, f"{timestamp}_{run_name}")
# Create directories
os.makedirs(run_dir, exist_ok=True)
os.makedirs(os.path.join(run_dir, "checkpoints"), exist_ok=True)
os.makedirs(os.path.join(run_dir, "logs"), exist_ok=True)
# Save run configuration
config_dict = vars(args)
with open(os.path.join(run_dir, "config.json"), 'w') as f:
json.dump(config_dict, f, indent=4)
return run_dir
def finetune():
run_dir = create_run_directory(args)
# Initialize wandb with the run directory
wandb_run_name = os.path.basename(run_dir)
wandb_run = wandb.init(
project="project_name",
config=args,
dir=os.path.join(run_dir, "logs")
)
# Save wandb run ID to a file
with open(os.path.join(run_dir, "wandb_run_id.txt"), "w") as f:
f.write(wandb_run.id)
# Create model and tokenizer
model, tokenizer = create_model_tokenizer_it(args)
# Data handling
train_dataset = load_and_preprocess_it(tokenizer=tokenizer, args=args)
data_collator = DataCollatorForSupervisedDataset(tokenizer=tokenizer)
data_module = dict(train_dataset=train_dataset, data_collator=data_collator)
model, abba_config = create_peft_model_it_abba(model, args)
param_counts = count_parameters(model, verbose=False)
total_params = param_counts['total_trainable_params']
classifier_params = param_counts['classifier_params']
non_classifier_params = param_counts['non_classifier_params']
wandb.log({"total_params": total_params, "classifier_params": classifier_params, "non_classifier_params": non_classifier_params})
# Setup optimizer
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr)
# Training arguments
training_args = TrainingArguments(
output_dir=os.path.join(run_dir, "checkpoints"),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
learning_rate=args.lr,
weight_decay=0,
warmup_ratio=args.warmup_ratio,
lr_scheduler_type=args.scheduler,
seed=args.seed,
report_to="wandb",
gradient_accumulation_steps=32,
save_strategy="no",
bf16=True,
tf32=False,
fp16=False,
logging_steps=1,
logging_first_step=True,
logging_dir=os.path.join(run_dir, "logs"),
)
# Save training arguments
training_args_path = os.path.join(run_dir, "training_args.json")
with open(training_args_path, 'w') as f:
json.dump(training_args.to_dict(), f, indent=4)
trainer = Trainer(
model=model,
args=training_args,
**data_module,
optimizers=(optimizer, None),
)
# Save tokenizer
tokenizer.save_pretrained(os.path.join(run_dir, "tokenizer"))
# Training
trainer.train()
# After training
final_model_path = os.path.join(run_dir, "final_model")
trainer.save_state()
model.save_pretrained(final_model_path)
return run_dir
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="LoRA SB for arithmetic reasoning tasks")
parser.add_argument("--data_path", type=str, default="meta-math/MetaMathQA", help="Path to the training data")
parser.add_argument("--dataset_split", type=str, default="train[:50000]", help="Dataset split to use. Options: ['train', 'test', 'eval']")
parser.add_argument("--dataset_field", type=str, nargs="+", default=["query", "response"], help="Fields of dataset input and output")
parser.add_argument("--model", type=str, default="mistralai/Mistral-7B-v0.1", help="Model name")
parser.add_argument("--lora_r", type=int, default=32, help="LoRA R value (assigns half of this to each adapter)")
parser.add_argument("--lora_alpha", type=int, default=16, help="LoRA alpha value")
parser.add_argument("--lora_dropout", type=float, default=0, help="LoRA dropout value")
parser.add_argument("--batch_size", type=int, default=1, help="Batch size")
parser.add_argument("--epochs", type=int, default=1, help="Number of epochs")
parser.add_argument("--scheduler", type=str, default="cosine", help="Learning rate scheduler")
parser.add_argument("--warmup_ratio", type=float, default=0.02, help="Warmup ratio")
parser.add_argument("--max_seq_length", type=int, default=512, help="Maximum sequence length")
parser.add_argument("--lr", type=float, default=1e-4, help="Learning rate")
parser.add_argument("--seed", type=int, default=42, help="Random seed")
parser.add_argument("--device", type=str, default="cuda", help="Device (cuda/cpu)")
args = parser.parse_args()
# Set random seeds
np.random.seed(args.seed)
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
# Run training
run_dir = finetune()