-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathutils.py
More file actions
276 lines (212 loc) · 8.83 KB
/
Copy pathutils.py
File metadata and controls
276 lines (212 loc) · 8.83 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
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
# most things here are taken from: https://github.com/Lightning-AI/litgpt/blob/main/litgpt/utils.py
from datetime import datetime
import math
import random
import os
import hydra
import numpy as np
from typing import Any, Iterable, List, Optional, Union
from typing_extensions import Self
from functools import partial
from tqdm import tqdm, trange
import torch
import torch.nn as nn
import torch.distributed as dist
from accelerate import Accelerator
from transformer import (
TransformerConfig,
Transformer,
)
from cat_transformer_adaptive import (
CAT_Config,
CAT_Transformer
)
def get_model(accelerate: Accelerator, cfg):
# pass hyperparameters from the yaml config file to the transformer config
if "transformer" == cfg.model.name:
config = TransformerConfig(
vocab_size=cfg.dataset.vocab_size,
block_size=cfg.model.block_size,
n_layer=cfg.model.n_layer,
dim=cfg.model.dim,
n_head=cfg.model.n_head,
norm_eps=cfg.train.norm_eps,
use_fused_ops=cfg.model.use_fused_ops,
use_qk_norm=cfg.model.use_qk_norm,
)
model = Transformer(config)
accelerate.print("Transformer config:", config)
accelerate.print(model)
return model
elif "cat_transformer" == cfg.model.name:
compressor_config = CAT_Config(
vocab_size=cfg.dataset.vocab_size,
block_size=cfg.model.block_size,
use_fused_ops=False, # liger-kernels doesn't support vmaps
use_qk_norm=cfg.model.use_qk_norm,
chunk_size=cfg.model.chunk_size,
dim=cfg.model.compressor_dim,
n_head=cfg.model.compressor_n_head,
dim_fx=cfg.model.dim_fx,
n_layer=cfg.model.compressor_n_layer,
) # layers are defined according to the paper, but one may use lower number of layers in the compressor
decoder_config = CAT_Config(
vocab_size=cfg.dataset.vocab_size,
block_size=cfg.model.block_size,
use_fused_ops=cfg.model.use_fused_ops,
use_qk_norm=cfg.model.use_qk_norm,
chunk_size=cfg.model.chunk_size,
dim=cfg.model.dim,
n_head=cfg.model.n_head,
n_layer=cfg.model.n_layer
)
model = CAT_Transformer(decoder_config, compressor_config)
accelerate.print("CAT compressor config:", compressor_config)
accelerate.print("CAT decoder config:", decoder_config)
accelerate.print(model)
return model
else:
raise ValueError(f"Unknown model type: {cfg['name']}")
# https://github.com/Lightning-AI/litgpt/blob/main/litgpt/pretrain.py#L384
@torch.no_grad()
def validate(accelerate: Accelerator, model: nn.Module, val_dataloader: torch.utils.data.DataLoader, cfg, chunk_size_powers=None):
print("Validating ...")
model.eval()
max_iters = cfg.eval.eval_iters
if chunk_size_powers is not None:
results = {}
for power in chunk_size_powers:
total_loss = 0.0
total_tokens = 0
desc = f"Evaluating (chunk={2**power})"
val_bar = tqdm(enumerate(val_dataloader), total=len(val_dataloader), desc=desc, disable=(not accelerate.is_main_process))
for k, batch in val_bar:
input_ids, targets = batch
if k >= max_iters:
break
input_ids, targets = input_ids.to(accelerate.device), targets.to(accelerate.device)
num_tokens = (targets != -100).sum().item()
with accelerate.autocast():
loss = model(input_ids, targets, chunk_size_power=power)
total_loss += loss.item() * num_tokens
total_tokens += num_tokens
val_bar.set_postfix_str(f"val loss: {total_loss / total_tokens:.4f}")
val_loss = total_loss / total_tokens
perplexity = math.exp(val_loss)
results[power] = (val_loss, perplexity)
accelerate.print(f" chunk_size={2**power}: loss={val_loss:.4f}, ppl={perplexity:.4f}")
model.train()
return results
total_loss = 0.0
total_tokens = 0
val_bar = tqdm(enumerate(val_dataloader), total=len(val_dataloader), desc="Evaluating", disable=(not accelerate.is_main_process))
for k, batch in val_bar:
input_ids, targets = batch
if k >= max_iters:
break
input_ids, targets = input_ids.to(accelerate.device), targets.to(accelerate.device)
num_tokens = (targets != -100).sum().item()
with accelerate.autocast():
loss = model(input_ids, targets)
total_loss += loss.item() * num_tokens
total_tokens += num_tokens
val_bar.set_postfix_str(f"val loss: {total_loss / total_tokens:.4f}")
val_loss = total_loss / total_tokens
perplexity = math.exp(val_loss)
model.train()
return val_loss, perplexity
# taken from: https://github.com/Lightning-AI/litgpt/blob/main/litgpt/pretrain.py#L299
# learning rate decay scheduler (cosine with linear warmup)
def get_lr(learning_rate: float, it: int, warmup_iters: int, max_iters: int, min_lr: float) -> float:
# 1) linear warmup for warmup_iters steps
if it < warmup_iters:
return learning_rate * it / warmup_iters
# 2) if it > max_iters, return min learning rate
if it > max_iters:
return min_lr
# 3) in between, use cosine decay down to min learning rate
decay_ratio = (it - warmup_iters) / (max_iters - warmup_iters)
assert 0 <= decay_ratio <= 1
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio)) # coeff ranges 0..1
return min_lr + coeff * (learning_rate - min_lr)
def num_parameters(module: nn.Module, requires_grad: Optional[bool] = None) -> int:
total = 0
for p in module.parameters():
if requires_grad is None or p.requires_grad == requires_grad:
if hasattr(p, "quant_state"):
# bitsandbytes 4bit layer support
total += math.prod(p.quant_state.shape)
else:
total += p.numel()
return total
class CycleIterator:
"""An iterator that cycles through an iterable indefinitely.
Example:
>>> iterator = CycleIterator([1, 2, 3])
>>> [next(iterator) for _ in range(5)]
[1, 2, 3, 1, 2]
Note:
Unlike ``itertools.cycle``, this iterator does not cache the values of the iterable.
"""
def __init__(self, iterable: Iterable, upper: Optional[int] = 999999) -> None:
self.iterable = iterable
self.epoch = 0
self.upper = upper
self.count = 0
self._iterator = None
def __next__(self) -> Any:
if self._iterator is None:
self._iterator = iter(self.iterable)
try:
if self.count >= self.upper:
self._iterator = iter(self.iterable)
self.count = 0
self.count += 1
return next(self._iterator)
except StopIteration:
self._iterator = iter(self.iterable)
self.epoch += 1
return next(self._iterator)
def __iter__(self) -> Self:
return self
def seed_everything(seed):
random.seed(seed)
os.environ['PYTHONHASHSEED'] = str(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # for multi-GPU
# # try to get deterministic training, but its slow!
# torch.backends.cudnn.deterministic = True
# torch.backends.cudnn.benchmark = False
def get_experiment_name(cfg, datetime_str, accelerate) -> str:
# start with todays date and time to make sure it is unique
name = cfg.wandb.exp_name
name += f" {datetime_str}"
return name
def create_results_dir(cfg, datetime_str, accelerate):
# first create a save dir name
name = f"{datetime_str}"
name = "/".join(name.split(" ")) # day/time
original_cwd = hydra.utils.get_original_cwd()
result_dir = os.path.join(original_cwd, cfg.results_dir, cfg.wandb.project, name)
os.makedirs(result_dir, exist_ok=True)
return result_dir
@torch.no_grad
def calculate_grad_norm(model, norm_type=2.0, scaler=None):
norm_type = float(norm_type)
grads = []
scale = scaler.get_scale() if scaler is not None else None
for p in model.parameters():
if p.grad is not None:
grad = p.grad.detach()
if scale is not None:
grad = grad / scale # No clone; just use the result of the division
grads.append(grad.float()) # Ensure float32 for stable norm computation
if not grads:
return 0.0
if norm_type == float("inf"):
total_norm = max(g.abs().max() for g in grads)
else:
total_norm = torch.norm(torch.stack([torch.norm(g, norm_type) for g in grads]), norm_type)
return total_norm.item()