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

Upload run_lora_chat.py with huggingface_hub

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