1earner1 commited on
Commit
2d1ec0d
·
verified ·
1 Parent(s): cc334e4

Upload run_lora_mistral_loop.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. run_lora_mistral_loop.py +579 -0
run_lora_mistral_loop.py ADDED
@@ -0,0 +1,579 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ model_name = "mistralai/Mistral-7B-Instruct-v0.3"
2
+ #model_name = "bigcode/starcoder2-7b"
3
+ #model_name = "dorkai/codeX-1.0" #"Alibaba-NLP/gte-Qwen1.5-7B-instruct" #"google/flan-t5-small" #"microsoft/Phi-3-medium-128k-instruct" #"google/gemma-2-9b-it" # "meta-llama/CodeLlama-7b-hf" #"deepseek-ai/DeepSeek-Coder-V2-Instruct" #"deepseek-ai/DeepSeek-Coder-V2-Lite-Instruct"
4
+ #out_name = "HPC_2_mistral_iffp_20k_5_lora" #"meta-llama/Meta-Llama-3-8B" #"tiiuae/falcon-40b" #"Phind/Phind-CodeLlama-34B-v2" # "deepseek-ai/DeepSeek-Coder-V2-Instruct" #
5
+ from datasets import load_dataset, Dataset
6
+ import pandas as pd
7
+ import json
8
+ import traceback
9
+ import peft
10
+ import os
11
+ from tqdm import tqdm
12
+ import sys
13
+ import math
14
+ from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, AdamW, default_data_collator, get_linear_schedule_with_warmup, get_cosine_schedule_with_warmup,set_seed
15
+ from torch.utils.data import DataLoader
16
+ import numpy as np
17
+ import os
18
+ import argparse
19
+ import torch
20
+ import datetime
21
+ from datasets import load_dataset
22
+ from transformers import (
23
+ AutoModelForCausalLM,
24
+ AutoTokenizer,
25
+ BitsAndBytesConfig,
26
+ HfArgumentParser,
27
+ TrainingArguments,
28
+ pipeline,
29
+ logging,
30
+ )
31
+ from peft import LoraConfig, PeftModel
32
+ from trl import SFTTrainer
33
+ import os
34
+ os.environ['WANDB_MODE'] = 'offline'
35
+ import wandb
36
+ import socket
37
+ import random
38
+
39
+ def set_seed(seed: int = 42):
40
+ random.seed(seed) # Python’s built-in random module
41
+ np.random.seed(seed) # NumPy
42
+ torch.manual_seed(seed) # PyTorch CPU
43
+ torch.cuda.manual_seed(seed) # PyTorch GPU
44
+ torch.cuda.manual_seed_all(seed) # If using multi-GPU
45
+ torch.backends.cudnn.deterministic = True # Ensures deterministic behavior in CuDNN
46
+ torch.backends.cudnn.benchmark = False # Disables benchmarking to maintain consistency
47
+
48
+ # Example usage
49
+ set_seed(42)
50
+ # import json
51
+ # filepath = "/kaggle/input/code-sim-try1/mutated_graph_all_lang_eq.json"
52
+ # examples = []
53
+ # with open(filepath, 'r') as file:
54
+ # for l in file:
55
+ # examples.append(json.loads(l))
56
+
57
+ # len(examples)
58
+ # examples[0]
59
+
60
+ # instruct_tune_dataset = load_dataset("mosaicml/instruct-v3",cache_dir = "/scratch/scai/mtech/aib222688/HF")
61
+ # instruct_tune_dataset = instruct_tune_dataset.filter(lambda x: x["source"] == "dolly_hhrlhf")
62
+
63
+
64
+ # traindataset_file = "./dataset_lfs/allpairs_data_large_900_loop2.json"
65
+ # valdataset_file = "./dataset_lfs/allpairs_data_val_900_loop2.json"
66
+ # testdataset_file = "./llm_for_code/datasets/codecontests/verified_iffp_900_loop2.json"
67
+
68
+ # initial_lr = 5e-6
69
+ # checkpoint_store_dir_path = "./HPC_2_mistral_iffp_20k_3_lora_900_loop2"
70
+
71
+ # num_epochs = 5
72
+ # batch_size_train = 1
73
+ # max_length = 2000
74
+ # ckpnt_NUM = 2000
75
+ # SAVEALL = False #True
76
+
77
+
78
+ parser = argparse.ArgumentParser(description='Run lora finetuning..., NOTE: UPDATE PEFT CONFIG if needed')
79
+ parser.add_argument('--model_name', default="mistralai/Mistral-7B-Instruct-v0.3",type=str)
80
+ parser.add_argument('--traindataset_files', nargs='+', type=str, default="./dataset_lfs/allpairs_data_large_900_loop2.json")
81
+ parser.add_argument('--valdataset_file', type=str, default="./dataset_lfs/allpairs_data_val_900_loop2.json")
82
+ parser.add_argument('--testdataset_file', type=str, default="./llm_for_code/datasets/codecontests/verified_iffp_900_loop2.json")
83
+ parser.add_argument('--checkpoint_store_dir_path', type=str, default="./HPC_3_mistral_iffp_20k_3_lora_900_loop2")
84
+ parser.add_argument('--initial_lr', type=float, default=5e-6)
85
+ parser.add_argument('--num_epochs', type=int, default=5)
86
+ parser.add_argument('--batch_size_train', type=int, default=1)
87
+ parser.add_argument('--ckpnt_num', type=int, default=2000)
88
+ parser.add_argument('--saveall', type=int, default=0)
89
+ parser.add_argument('--prompt_file_path', type=str, default='./loop_prompt.txt')
90
+ parser.add_argument('--max_length', type=int, default=2000)
91
+ parser.add_argument('--max_new_tok', type=int, default=50)
92
+
93
+ args = parser.parse_args()
94
+ print(f"{len(vars(args))=}")
95
+
96
+
97
+
98
+
99
+
100
+ model_name = args.model_name.replace('\r', '')
101
+ traindataset_files = args.traindataset_files
102
+ for i in range(len(traindataset_files)):
103
+ traindataset_files[i] = traindataset_files[i].replace('\r', '')
104
+ valdataset_file = args.valdataset_file.replace('\r', '')
105
+ testdataset_file = args.testdataset_file.replace('\r', '')
106
+ checkpoint_store_dir_path = args.checkpoint_store_dir_path.replace('\r', '')
107
+ initial_lr = args.initial_lr
108
+ num_epochs = args.num_epochs
109
+ batch_size_train = args.batch_size_train
110
+ ckpnt_NUM = args.ckpnt_num
111
+ SAVEALL = args.saveall
112
+ prompt_file_path = args.prompt_file_path.replace('\r', '')
113
+ max_length = args.max_length
114
+ max_new_tok = args.max_new_tok
115
+
116
+
117
+
118
+ hostname = socket.gethostname()
119
+ ip_address = socket.gethostbyname(hostname)
120
+ node_name = os.uname().nodename
121
+ system_info = os.uname()
122
+ machine_info = {
123
+ "hostname": hostname,
124
+ "ip_address": ip_address,
125
+ "node_name": node_name,
126
+ "system_info": {
127
+ "sysname": system_info.sysname,
128
+ "nodename": system_info.nodename,
129
+ "release": system_info.release,
130
+ "version": system_info.version,
131
+ "machine": system_info.machine,
132
+ },
133
+ }
134
+
135
+ os.makedirs(checkpoint_store_dir_path, exist_ok=True)
136
+ current_time = datetime.datetime.now()
137
+
138
+ with open(checkpoint_store_dir_path+'/lora_logs.txt', 'a') as log_file:
139
+ log_file.write(f"{current_time}: running lora\n {vars(args)}\n")
140
+ log_file.write(f"{machine_info}-----\n")
141
+
142
+ traindata = []
143
+ nf4_config = BitsAndBytesConfig(
144
+ load_in_4bit=True,
145
+ bnb_4bit_quant_type="nf4",
146
+ bnb_4bit_use_double_quant=True,
147
+ bnb_4bit_compute_dtype=torch.bfloat16
148
+ )
149
+
150
+ #mpath = './codellama'
151
+ # model = AutoModelForCausalLM.from_pretrained(
152
+ # model_name,
153
+ # #device_map='auto',
154
+ # #quantization_config=nf4_config,
155
+ # use_cache=True,
156
+ # #cache_dir = "../aib222688.scratch/HF/",
157
+ # attn_implementation="sdpa", #"flash_attention_2",
158
+ # torch_dtype=torch.float16,
159
+ # #trust_remote_code=True,
160
+ # )
161
+ model = AutoModelForCausalLM.from_pretrained(
162
+ #mpath,
163
+ model_name,
164
+ #use_cache=True,
165
+ #cache_dir = "../aib222688.scratch/HF/",
166
+ #attn_implementation="flash_attention_2",
167
+ torch_dtype=torch.float16,
168
+ #device_map='auto',
169
+ #quantization_config=nf4_config,
170
+ #use_cache=False
171
+ )
172
+
173
+ print(f"Shards loaded for {model_name}")
174
+ # for name, module in model.named_modules():
175
+ # print(f"{name}: {module}")
176
+
177
+ # model = AutoModelForCausalLM.from_pretrained(
178
+ # "./HPC_2_mistral_iffp_20k_2_lora_1200/checkpoint_0_18000/"
179
+ # )
180
+
181
+ tokenizer = AutoTokenizer.from_pretrained(model_name)
182
+ #tokenizer = AutoTokenizer.from_pretrained(mpath)
183
+
184
+ tokenizer.pad_token = tokenizer.eos_token
185
+ tokenizer.padding_side = "right"
186
+
187
+
188
+ #wandb.login(key='e7fdeef2a423ceed55ae12d2c9f1bc530a9e9331')
189
+ wandb.init(project=checkpoint_store_dir_path[3:]+'_wandb', config={
190
+ 'args' : vars(args),
191
+ 'machine' : machine_info,
192
+ })
193
+
194
+
195
+ for traindataset_file in traindataset_files:
196
+ with open(traindataset_file, 'r') as file:
197
+ for l in file:
198
+ traindata.append(json.loads(l))
199
+
200
+ valdata = []
201
+
202
+ with open(valdataset_file, 'r') as file:
203
+ for l in file:
204
+ valdata.append(json.loads(l))
205
+
206
+ testdata = []
207
+
208
+
209
+ with open(testdataset_file, 'r') as file:
210
+ for l in file:
211
+ testdata.append(json.loads(l))
212
+ print(len(traindata), len(valdata), len(testdata))
213
+
214
+
215
+
216
+
217
+
218
+ file_content = ""
219
+ with open(prompt_file_path, 'r') as file:
220
+ file_content = file.read()
221
+
222
+
223
+ def create_prompt(pair):
224
+ bos_token = "<s>"
225
+ eos_token = "</s>"
226
+ if pair['prog1']['probid'] == pair['prog2']['probid']:
227
+ response = "Yes"
228
+ else:
229
+ response = "No"
230
+
231
+ full_prompt = ""
232
+ full_prompt += bos_token
233
+ #print(f"{pair['prog1']['scode']=}")
234
+ full_prompt += file_content + pair['prog1']['scode'] + "\nProgram 2:"
235
+ full_prompt += pair['prog2']['scode']+ "\n### Response:"
236
+ full_prompt += "\n" #+ response
237
+ full_prompt += eos_token
238
+
239
+ return full_prompt, response
240
+ #print(create_prompt(instruct_tune_dataset["train"][1]))
241
+
242
+ traindata1 = list(traindata) #[0:]
243
+ valdata1 = list( valdata)
244
+ testdata1 = list(testdata) #[0:]
245
+
246
+ traindata = []
247
+ pos_cnt = 0
248
+ for tdata in traindata1:
249
+ inp, trg = create_prompt(tdata)
250
+ if tdata['prog1']['probid'] == tdata['prog2']['probid']: # if pair['label'] == 1: #
251
+ pos_cnt += 1
252
+ traindata.append({
253
+ 'inputs' : inp,
254
+ 'targets' : trg
255
+ })
256
+ print("traindata[0] ", traindata[0], pos_cnt)
257
+
258
+ valdata = []
259
+
260
+ for tdata in valdata1:
261
+ inp, trg = create_prompt(tdata)
262
+ valdata.append({
263
+ 'inputs' : inp,
264
+ 'targets' : trg
265
+ })
266
+
267
+ testdata = []
268
+
269
+ for tdata in testdata1:
270
+ inp, trg = create_prompt(tdata)
271
+ testdata.append({
272
+ 'inputs' : inp,
273
+ 'targets' : trg
274
+ })
275
+
276
+ traindataset = Dataset.from_pandas(pd.DataFrame(traindata))
277
+ valdataset = Dataset.from_pandas(pd.DataFrame(valdata))
278
+ testdataset = Dataset.from_pandas(pd.DataFrame(testdata))
279
+
280
+ instruct_tune_dataset = {"train": traindataset,
281
+ "val" : valdataset,
282
+ "test" : testdataset}
283
+
284
+
285
+
286
+
287
+
288
+
289
+
290
+
291
+
292
+ def preprocess_function(examples):
293
+ batch_size = len(examples['inputs'])
294
+ #inputs = [f"<s>[INST] Question : {x} [/INST] \\n Answer : " for x in examples[past_context_code]]
295
+ #inputs = [f"\n<|user|>\n You are given a set of APIs and previously generated Code as context. The task is given a new requirement from Bob modify or expand the given code using the provided APIs.\n\nAPIs:\n{apis}\n\nContext:\n{past_context}\n\nInput:\n{new_input} \n<|assistant|>\n " for apis, past_context, new_input in zip(examples['apis'], examples['past_context_code'], examples['new_input'])]
296
+ #targets = [str(x) for x in examples[label_column]]
297
+ #inputs, targets = get_examples_all_context(examples)
298
+
299
+ #inputs, targets = get_examples_all_context(examples)
300
+ #inputs, targets = get_examples_all_context_granite(examples, only_code=False)
301
+ inputs = []
302
+ targets = []
303
+ # for eg in examples:
304
+ # print(eg)
305
+ # #inp, trg = create_prompt(eg)
306
+ # #inputs.append(inp)
307
+ # #targets.append(trg)
308
+ inputs = examples['inputs']
309
+ targets = examples['targets']
310
+ model_inputs = tokenizer(inputs)
311
+
312
+ #print("Input example:\n{}".format(inputs[0]))
313
+ #print("Output example:\n{}".format(targets[0]))
314
+
315
+ input_sizes = [len(tokens) for tokens in model_inputs['input_ids']]
316
+ #print("Input sizes {}".format(input_sizes))
317
+ labels = tokenizer(targets, add_special_tokens=False) # don't add bos token because we concatenate with inputs
318
+ label_sizes = [len(tokens) for tokens in labels['input_ids']]
319
+ #print("Label sizes {}".format(label_sizes))
320
+
321
+ for i in range(batch_size):
322
+ sample_input_ids = model_inputs["input_ids"][i]
323
+ label_input_ids = labels["input_ids"][i] + [tokenizer.eos_token_id]
324
+ # print(i, sample_input_ids, label_input_ids)
325
+ model_inputs["input_ids"][i] = sample_input_ids + label_input_ids
326
+ labels["input_ids"][i] = [-100] * len(sample_input_ids) + label_input_ids
327
+ model_inputs["attention_mask"][i] = [1] * len(model_inputs["input_ids"][i])
328
+ # print(model_inputs)
329
+ for i in range(batch_size):
330
+ sample_input_ids = model_inputs["input_ids"][i]
331
+ label_input_ids = labels["input_ids"][i]
332
+ model_inputs["input_ids"][i] = [tokenizer.pad_token_id] * (
333
+ max_length - len(sample_input_ids)
334
+ ) + sample_input_ids
335
+ model_inputs["attention_mask"][i] = [0] * (max_length - len(sample_input_ids)) + model_inputs["attention_mask"][i]
336
+ labels["input_ids"][i] = [-100] * (max_length - len(sample_input_ids)) + label_input_ids
337
+ model_inputs["input_ids"][i] = torch.tensor(model_inputs["input_ids"][i][:max_length])
338
+ model_inputs["attention_mask"][i] = torch.tensor(model_inputs["attention_mask"][i][:max_length])
339
+ labels["input_ids"][i] = torch.tensor(labels["input_ids"][i][:max_length])
340
+ model_inputs["labels"] = labels["input_ids"]
341
+ input_sizes = [len(tokens) for tokens in model_inputs['input_ids']]
342
+ #print("Input sizes {}".format(input_sizes))
343
+ return model_inputs
344
+
345
+ processed_datasets = traindataset.map(
346
+ preprocess_function,
347
+ batched=True,
348
+ num_proc=1,
349
+ remove_columns=traindataset.column_names,
350
+ load_from_cache_file=False,
351
+ desc="Running tokenizer on dataset",
352
+ )
353
+ train_dataset = processed_datasets
354
+ train_dataloader = DataLoader(
355
+ train_dataset, shuffle=True, collate_fn=default_data_collator, batch_size=batch_size_train, pin_memory=True
356
+ )
357
+
358
+
359
+ processed_datasets = valdataset.map(
360
+ preprocess_function,
361
+ batched=True,
362
+ num_proc=1,
363
+ remove_columns=valdataset.column_names,
364
+ load_from_cache_file=False,
365
+ desc="Running tokenizer on dataset",
366
+ )
367
+ val_dataset = processed_datasets
368
+ val_dataloader = DataLoader(
369
+ val_dataset, shuffle=True, collate_fn=default_data_collator, batch_size=batch_size_train, pin_memory=True
370
+ )
371
+
372
+
373
+
374
+
375
+ def test_preprocess_function(examples):
376
+ batch_size = len(examples['inputs'])
377
+ #inputs, targets = get_examples_all_context(examples)
378
+ #inputs, targets = get_examples_all_context_granite(examples)
379
+ #inputs = [f"\n<|user|>\n You are given a set of APIs and previously generated Code as context. The task is given a new requirement from Bob modify or expand the given code using the provided APIs.\n\nAPIs:\n{apis}\n\nContext:\n{past_context}\n\nInput:\n{new_input} \n<|assistant|>\n " for apis, past_context, new_input in zip(examples['apis'], examples['past_context_code'], examples['new_input'])]
380
+ model_inputs = tokenizer(examples['inputs'])
381
+ # print(model_inputs)
382
+ for i in range(batch_size):
383
+ sample_input_ids = model_inputs["input_ids"][i]
384
+ model_inputs["input_ids"][i] = [tokenizer.pad_token_id] * (
385
+ max_length - len(sample_input_ids)
386
+ ) + sample_input_ids
387
+ model_inputs["attention_mask"][i] = [0] * (max_length - len(sample_input_ids)) + model_inputs["attention_mask"][i]
388
+ model_inputs["input_ids"][i] = torch.tensor(model_inputs["input_ids"][i][:max_length])
389
+ model_inputs["attention_mask"][i] = torch.tensor(model_inputs["attention_mask"][i][:max_length])
390
+ return model_inputs
391
+
392
+
393
+ processed_datasets = testdataset.map(
394
+ preprocess_function,
395
+ batched=True,
396
+ num_proc=1,
397
+ remove_columns=testdataset.column_names,
398
+ load_from_cache_file=False,
399
+ desc="Running tokenizer on dataset",
400
+ )
401
+ test_dataset = processed_datasets
402
+ test_dataloader = DataLoader(
403
+ test_dataset, shuffle=False, collate_fn=default_data_collator, batch_size=batch_size_train, pin_memory=True
404
+ )
405
+
406
+
407
+
408
+
409
+ peft_config = LoraConfig(
410
+ lora_alpha=16,
411
+ lora_dropout=0.1,
412
+ #target_modules = ['c_attn'],
413
+ target_modules = ['q_proj', 'k_proj', 'v_proj', 'o_proj'], #qwen
414
+ #target_modules = ['q_proj', 'v_proj'], #mistral
415
+ r=64,
416
+ bias="none",
417
+ task_type="CAUSAL_LM"
418
+ )
419
+
420
+ # peft_config = LoraConfig(
421
+ # r=lora_r,
422
+ # lora_alpha=lora_alpha,
423
+ # lora_dropout=lora_dropout,
424
+ # target_modules= target_modules,
425
+ # bias="none",
426
+ # task_type="CAUSAL_LM"
427
+ # )
428
+
429
+ model = peft.get_peft_model(model, peft_config)
430
+
431
+ wandb.watch(model, log='all')
432
+ print("Model loaded successfully!")
433
+
434
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
435
+
436
+ #optimizer = AdamW(model.parameters(), lr=3e-4)
437
+ optimizer = AdamW(model.parameters(), lr=initial_lr)
438
+
439
+ # Instantiate scheduler
440
+ lr_scheduler = get_cosine_schedule_with_warmup(
441
+ optimizer=optimizer,
442
+ num_warmup_steps=0.06 * (len(train_dataloader) * num_epochs),
443
+ num_training_steps=(len(train_dataloader) * num_epochs),
444
+ )
445
+
446
+ model.to(device)
447
+ model.to('cuda')
448
+
449
+ #model = torch.nn.DataParallel(model)
450
+ #model = model.cuda()
451
+ the_best_eval_loss = 10000
452
+ for epoch in range(num_epochs):
453
+ try:
454
+ model.train()
455
+ total_loss = 0
456
+ best_eval_loss = 10000 #np.inf
457
+
458
+ for step, batch in enumerate(tqdm(train_dataloader)):
459
+ batch = {k: v.to(device) for k, v in batch.items()}
460
+ #batch = {k: v.cuda() for k, v in batch.items()}
461
+ # print(batch)
462
+ #print(batch["input_ids"].shape)
463
+ # if step > 5:
464
+ # break
465
+ #batch.to(device)
466
+ outputs = model(**batch)
467
+ loss = outputs.loss
468
+ total_loss += loss.detach().float()
469
+ wandb.log({'train_loss': loss})
470
+ wandb.log({'lr': lr_scheduler.get_last_lr()[0], 'step': step})
471
+ if step % 100 == 0:
472
+ wandb.log({'train_step_loss': loss})
473
+ print(loss)
474
+ loss.backward()
475
+ #print("Loss {}".format(loss.item()))
476
+ optimizer.step()
477
+ lr_scheduler.step()
478
+ optimizer.zero_grad()
479
+ # if step % ckpnt_NUM == 0:
480
+ # checkpoint_dir = os.path.join(checkpoint_store_dir_path, f"checkpoint_{epoch}_{step}/")
481
+ # os.makedirs(checkpoint_dir, exist_ok=True)
482
+ # model.save_pretrained(checkpoint_dir)
483
+
484
+ if step % ckpnt_NUM == 0:
485
+
486
+ model.eval()
487
+ eval_loss = 0
488
+ eval_preds = []
489
+ eval_cnt = 1
490
+ for step1, batch_eval in enumerate(tqdm(val_dataloader)):
491
+
492
+ batch_eval = {k: v.to(device) for k, v in batch_eval.items()}
493
+ #batch_eval = {k: v.cuda() for k, v in batch_eval.items()}
494
+
495
+ #outputs = model.generate(**batch_eval, max_new_tokens=48)
496
+ #out = tokenizer.batch_decode(outputs, skip_special_tokens=True)
497
+ # for x in out:
498
+ # print(x)
499
+ # print("#" * 50)
500
+ with torch.no_grad():
501
+ outputs = model(**batch_eval)
502
+ loss = outputs.loss
503
+ if not math.isnan(loss.detach().float()) :
504
+ eval_loss += loss.detach().float()
505
+ eval_cnt += 1
506
+ # eval_preds.extend(
507
+ # tokenizer.batch_decode(torch.argmax(outputs.logits, -1).detach().cpu().numpy(),
508
+ # skip_special_tokens=True)
509
+ # )
510
+ print(eval_loss)
511
+ wandb.log({'eval_loss': eval_loss})
512
+ eval_epoch_loss = eval_loss / len(val_dataloader)
513
+ if ((eval_loss/eval_cnt) < best_eval_loss) or SAVEALL==1:
514
+
515
+ best_eval_loss = eval_loss/eval_cnt
516
+ print(f"saving...{best_eval_loss} to checkpoint_{epoch}_{step}\n")
517
+ checkpoint_dir = os.path.join(checkpoint_store_dir_path, f"checkpoint_{epoch}_{step}/")
518
+ os.makedirs(checkpoint_dir, exist_ok=True)
519
+ model.save_pretrained(checkpoint_dir)
520
+
521
+ if ((eval_loss/eval_cnt) < the_best_eval_loss):
522
+ the_best_eval_loss = eval_loss/eval_cnt
523
+ print(f"saving...{best_eval_loss} to checkpoint_{epoch}_{step} is best so far\n")
524
+ wandb.log({'thebest_loss': epoch, 'thebest_step' : step})
525
+
526
+
527
+ eval_ppl = torch.exp(eval_epoch_loss)
528
+ train_epoch_loss = total_loss #/ len(train_dataloader)
529
+ train_ppl = torch.exp(train_epoch_loss)
530
+ print(f"{epoch=}: {train_ppl=} {train_epoch_loss=} {eval_ppl=} {eval_epoch_loss=} {eval_loss/eval_cnt=} {eval_cnt=}")
531
+ #print(f"{epoch=}: {train_ppl=} {train_epoch_loss=}")
532
+
533
+ print("Total Loss {}".format(total_loss.item()))
534
+ if ((epoch+1) % 1) == 0:
535
+ checkpoint_dir = os.path.join(checkpoint_store_dir_path, f"checkpoint_{epoch}/")
536
+ os.makedirs(checkpoint_dir, exist_ok=True)
537
+ model.save_pretrained(checkpoint_dir)
538
+ model.eval()
539
+ eval_loss = 0
540
+ eval_preds = []
541
+ for step1, batch_eval in enumerate(tqdm(test_dataloader)):
542
+ if step1 > 5:
543
+ break
544
+ batch_eval = {k: v.to(device) for k, v in batch_eval.items()}
545
+ #batch_eval = {k: v.cuda() for k, v in batch_eval.items()}
546
+
547
+ outputs = model.generate(**batch_eval, max_new_tokens=max_new_tok)
548
+ out = tokenizer.batch_decode(outputs, skip_special_tokens=True)
549
+ for x in out:
550
+ print(x)
551
+ print("#" * 50)
552
+ # with torch.no_grad():
553
+ # outputs = model(**batch)
554
+ # loss = outputs.loss
555
+ # eval_loss += loss.detach().float()
556
+ # eval_preds.extend(
557
+ # tokenizer.batch_decode(torch.argmax(outputs.logits, -1).detach().cpu().numpy(),
558
+ # skip_special_tokens=True)
559
+ # )
560
+
561
+ # eval_epoch_loss = eval_loss / len(test_dataloader)
562
+ # eval_ppl = torch.exp(eval_epoch_loss)
563
+ train_epoch_loss = total_loss / len(train_dataloader)
564
+ train_ppl = torch.exp(train_epoch_loss)
565
+ #print(f"{epoch=}: {train_ppl=} {train_epoch_loss=} {eval_ppl=} {eval_epoch_loss=}")
566
+ print(f"{epoch=}: {train_ppl=} {train_epoch_loss=}")
567
+
568
+
569
+
570
+
571
+ except KeyboardInterrupt:
572
+
573
+ checkpoint_dir = os.path.join(checkpoint_store_dir_path, f"checkpoint_{epoch}_interrupt/")
574
+ os.makedirs(checkpoint_dir, exist_ok=True)
575
+ model.save_pretrained(checkpoint_dir)
576
+
577
+
578
+
579
+