-
Notifications
You must be signed in to change notification settings - Fork 23
Expand file tree
/
Copy pathtrainer.py
More file actions
108 lines (85 loc) · 3.6 KB
/
Copy pathtrainer.py
File metadata and controls
108 lines (85 loc) · 3.6 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
import os
import gc
import torch
import wandb
from tqdm import tqdm
from abc import ABCMeta, abstractmethod
class Trainer(metaclass=ABCMeta):
def __init__(self, save_root, accel):
self.save_root = save_root
self.accel = accel
@abstractmethod
def compute_loss_and_update(self):
pass
def memory_optimization(self):
# wait
self.accel.wait_for_everyone()
# memory deallocation
gc.collect()
# removing cache
torch.cuda.empty_cache()
# wait
self.accel.wait_for_everyone()
def wait_and_save_ckpt(self, **kwargs):
model = kwargs['model']
batch_ind = kwargs['batch_ind']
length_dataloader = kwargs['epochs'] * len(kwargs['train_dataloader'])
save_number = kwargs['save_number']
processor = kwargs['processor']
# wait for everyone
self.accel.wait_for_everyone()
if batch_ind+1 in [int(i/save_number*length_dataloader) for i in range(1, save_number+1)]:
# Student
unwrapped_model = self.accel.unwrap_model(model)
unwrapped_model.save_pretrained(
os.path.join(self.save_root, f'{batch_ind+1}'),
is_main_process=self.accel.is_main_process and self.accel.local_process_index==0,
save_function=self.accel.save,
state_dict=self.accel.get_state_dict(model),
max_shard_size='3GB'
)
# processor
processor.save_pretrained(os.path.join(self.save_root, f'{batch_ind+1}'))
# print
self.accel.print(f"----{batch_ind+1}: Save Comleted!!----")
# wait for everyone
self.accel.wait_for_everyone()
def train(self, **kwargs):
"""
necessary kwargs
- model
- vllm_model
- epochs
- train_dataloader
- optimizer
- scheduler
- processor
- max_new_tokens
- wandb
- save_number
"""
for epoch in range(kwargs['epochs']):
# progress bar
prog_bar = tqdm(enumerate(kwargs['train_dataloader']),
disable=not (self.accel.is_main_process and self.accel.local_process_index==0),
total=len(kwargs['train_dataloader']))
# training start
for batch_ind, inputs in prog_bar:
# memory opt
self.memory_optimization()
# forward & backward
with self.accel.accumulate(kwargs['model']):
# backwarding loss with gradient accumulation
loss_dict = self.compute_loss_and_update(inputs, **kwargs)
# wandb logging
if kwargs['wandb'] and self.accel.is_main_process and self.accel.local_process_index==0:
update_wandb_dict = {'lr': kwargs['scheduler'].get_last_lr()[0]}
for k, v in loss_dict.items(): update_wandb_dict.update({k: v})
wandb.log(update_wandb_dict)
# displaying progress bar
GPU0_usage = torch.cuda.memory_reserved(device=0) / 1024**3
prog_bar.set_description(f"[GPU0:{GPU0_usage:.0f}][LR:{kwargs['scheduler'].get_last_lr()[0]:.6f}] " +\
" | ".join([f"{k}: {v:.3f}" for k, v in loss_dict.items()]), refresh=True)
# saving the model
kwargs['batch_ind'] = epoch * len(kwargs['train_dataloader']) + batch_ind
self.wait_and_save_ckpt(**kwargs)