LMFuser 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- lmfuser/__init__.py +0 -0
- lmfuser/model_loader.py +33 -0
- lmfuser/optimizers.py +162 -0
- lmfuser/runners/__init__.py +1 -0
- lmfuser/runners/ddp_runner.py +619 -0
- lmfuser/runners/runner.py +40 -0
- lmfuser/schedulers.py +210 -0
- lmfuser/task.py +287 -0
- lmfuser/utils.py +249 -0
- lmfuser-0.0.1.dist-info/METADATA +30 -0
- lmfuser-0.0.1.dist-info/RECORD +13 -0
- lmfuser-0.0.1.dist-info/WHEEL +4 -0
- lmfuser-0.0.1.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,619 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
from typing import Any, Iterator, TypeVar, Generic, Union, Hashable, List, Iterable
|
|
3
|
+
from contextlib import ExitStack
|
|
4
|
+
from collections import defaultdict
|
|
5
|
+
import logging
|
|
6
|
+
import os
|
|
7
|
+
import random
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
import json
|
|
10
|
+
from logging import Logger, getLogger
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
from torch import nn
|
|
14
|
+
import torch.distributed
|
|
15
|
+
from torch.nn.parallel import DistributedDataParallel as DDP
|
|
16
|
+
try:
|
|
17
|
+
from torch.amp.grad_scaler import GradScaler
|
|
18
|
+
OLD_GRADSCALER = False
|
|
19
|
+
except ImportError:
|
|
20
|
+
from torch.cuda.amp.grad_scaler import GradScaler
|
|
21
|
+
OLD_GRADSCALER = True
|
|
22
|
+
from torch.amp.autocast_mode import autocast
|
|
23
|
+
from torch.nn.utils.clip_grad import clip_grad_norm_
|
|
24
|
+
from torch.optim import Optimizer
|
|
25
|
+
from torch import Tensor
|
|
26
|
+
from torch.optim.lr_scheduler import LRScheduler
|
|
27
|
+
from tqdm import tqdm
|
|
28
|
+
from hyperargs import Conf, StrArg, IntArg, FloatArg, OptionArg, BoolArg
|
|
29
|
+
from lmfuser_data.interfaces import Batch
|
|
30
|
+
import wandb
|
|
31
|
+
from wandb.wandb_run import Run
|
|
32
|
+
|
|
33
|
+
from ..task import Task, Tasks
|
|
34
|
+
from ..utils import (
|
|
35
|
+
get_global_rank,
|
|
36
|
+
get_local_rank,
|
|
37
|
+
get_world_size,
|
|
38
|
+
get_default_device,
|
|
39
|
+
dist_init,
|
|
40
|
+
dist_avg,
|
|
41
|
+
batch_all_gather,
|
|
42
|
+
gather_object,
|
|
43
|
+
cal_acc_num,
|
|
44
|
+
get_default_device_type
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
from ..optimizers import OptimizerConfig
|
|
48
|
+
from ..schedulers import LRSchedulerConfig
|
|
49
|
+
from ..model_loader import ModelLoader, ModelLoaderConf
|
|
50
|
+
from .runner import RunerConf, Runner
|
|
51
|
+
|
|
52
|
+
logger = logging.getLogger(__name__)
|
|
53
|
+
|
|
54
|
+
T = TypeVar('T', bound=Conf)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class Wrapper(nn.Module):
|
|
58
|
+
|
|
59
|
+
def __init__(self, model: nn.Module) -> None:
|
|
60
|
+
super().__init__()
|
|
61
|
+
self.module = model.to(get_default_device())
|
|
62
|
+
self.forward = model.forward
|
|
63
|
+
|
|
64
|
+
def __getattr__(self, name: str) -> Any:
|
|
65
|
+
try:
|
|
66
|
+
return super().__getattr__(name)
|
|
67
|
+
except:
|
|
68
|
+
return self.module.__getattr__(name)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class DDPWraper(nn.Module):
|
|
72
|
+
|
|
73
|
+
def __init__(self, model: nn.Module) -> None:
|
|
74
|
+
super().__init__()
|
|
75
|
+
self.model = DDP(
|
|
76
|
+
model.to(get_default_device()), find_unused_parameters=True
|
|
77
|
+
) if get_world_size() > 1 else Wrapper(model)
|
|
78
|
+
self.forward = self.model.forward
|
|
79
|
+
|
|
80
|
+
def __getattr__(self, name: str) -> Any:
|
|
81
|
+
try:
|
|
82
|
+
return super().__getattr__(name)
|
|
83
|
+
except:
|
|
84
|
+
return self.model.module.__getattr__(name)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class DDPRunnerConfig(RunerConf):
|
|
88
|
+
|
|
89
|
+
checkpoint_directory = StrArg('please set checkpoint directory here!')
|
|
90
|
+
model_loader_conf = ModelLoaderConf()
|
|
91
|
+
|
|
92
|
+
stop_by = OptionArg('step', options=['step', 'epoch'])
|
|
93
|
+
total_step = IntArg(100, min_value=1)
|
|
94
|
+
total_epoch = IntArg(10, min_value=1)
|
|
95
|
+
eval_step_freq = IntArg(100, min_value=1)
|
|
96
|
+
save_step_freq = IntArg(100, min_value=1)
|
|
97
|
+
|
|
98
|
+
batch_size = IntArg(32, min_value=1)
|
|
99
|
+
sub_batch_size = IntArg(32, min_value=1)
|
|
100
|
+
|
|
101
|
+
grad_norm_clip = FloatArg(None, min_value=0.0, allow_none=True)
|
|
102
|
+
|
|
103
|
+
task_conf = Tasks()
|
|
104
|
+
|
|
105
|
+
optimizer: OptimizerConfig = OptimizerConfig()
|
|
106
|
+
lr_scheduler: LRSchedulerConfig = LRSchedulerConfig()
|
|
107
|
+
|
|
108
|
+
dp_type = OptionArg(default='ddp', options=['ddp'])
|
|
109
|
+
model_precision = OptionArg(options=['fp32', 'fp16'], default='fp32')
|
|
110
|
+
use_amp = BoolArg(default=False)
|
|
111
|
+
amp_precision = OptionArg(options=['fp16', 'bf16'], default='fp16')
|
|
112
|
+
seed = IntArg(42)
|
|
113
|
+
|
|
114
|
+
ignore_data_error = BoolArg(default=False)
|
|
115
|
+
data_row_qps = FloatArg(None, min_value=0.0, allow_none=True)
|
|
116
|
+
instruct_timeout = FloatArg(30.0, min_value=0.0)
|
|
117
|
+
worker_timeout = FloatArg(30.0, min_value=0.0)
|
|
118
|
+
shuffle_dataset = BoolArg(default=True)
|
|
119
|
+
row_prefetch = IntArg(0, min_value=0)
|
|
120
|
+
num_row_workers = IntArg(1, min_value=1)
|
|
121
|
+
|
|
122
|
+
resume_training = BoolArg(default=False)
|
|
123
|
+
resume_path = StrArg(default=None, allow_none=True)
|
|
124
|
+
|
|
125
|
+
@property
|
|
126
|
+
def _default_precision(self) -> torch.dtype:
|
|
127
|
+
if self.model_precision == 'fp32':
|
|
128
|
+
return torch.float32
|
|
129
|
+
elif self.model_precision == 'fp16':
|
|
130
|
+
return torch.float16
|
|
131
|
+
else:
|
|
132
|
+
raise ValueError(self.model_precision)
|
|
133
|
+
|
|
134
|
+
@property
|
|
135
|
+
def _num_acc_steps(self) -> int:
|
|
136
|
+
bs = self.batch_size.value()
|
|
137
|
+
sbs = self.sub_batch_size.value()
|
|
138
|
+
assert bs is not None and sbs is not None
|
|
139
|
+
|
|
140
|
+
if bs % (sbs * get_world_size()) != 0:
|
|
141
|
+
raise ValueError(
|
|
142
|
+
f'batch_size ({bs}) must be divisible by sub_batch_size * world_size ({sbs * get_world_size()})'
|
|
143
|
+
)
|
|
144
|
+
return bs // (sbs * get_world_size())
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
class DDPRunner(Runner[DDPRunnerConfig]):
|
|
148
|
+
|
|
149
|
+
def __init__(self, config: DDPRunnerConfig, *args, **kwargs) -> None:
|
|
150
|
+
super().__init__(config, *args, **kwargs)
|
|
151
|
+
if get_world_size() > 1:
|
|
152
|
+
dist_init()
|
|
153
|
+
|
|
154
|
+
self.tasks = [task.conf for task in config.task_conf.tasks]
|
|
155
|
+
self.step = 1
|
|
156
|
+
self.pre_epoch = 0
|
|
157
|
+
self.model_loader = config.model_loader_conf.get_model_loader()
|
|
158
|
+
|
|
159
|
+
if config.resume_training.value():
|
|
160
|
+
resume_path = config.resume_path.value()
|
|
161
|
+
assert resume_path is not None, 'resume_path is None'
|
|
162
|
+
self.load(resume_path)
|
|
163
|
+
self.config.seed = self.config.seed.parse(hash(f'original_seed_{self.config.seed.value()}|step_{self.step}'))
|
|
164
|
+
|
|
165
|
+
self.train_data_loaders = config.task_conf.get_train_dataloaders(
|
|
166
|
+
batch_size=config.sub_batch_size.value(), # type: ignore
|
|
167
|
+
seed=config.seed.value(), # type: ignore
|
|
168
|
+
shuffle=config.shuffle_dataset.value(), # type: ignore
|
|
169
|
+
prefetch_factor=config.row_prefetch.value(), # type: ignore
|
|
170
|
+
num_workers=config.num_row_workers.value(), # type: ignore
|
|
171
|
+
ignore_error=config.ignore_data_error.value(), # type: ignore
|
|
172
|
+
qps=config.data_row_qps.value(),
|
|
173
|
+
instruct_timeout=config.instruct_timeout.value(), # type: ignore
|
|
174
|
+
worker_timeout=config.worker_timeout.value(), # type: ignore
|
|
175
|
+
world_size=get_world_size(),
|
|
176
|
+
rank=get_global_rank(),
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
self.eval_data_loaders = config.task_conf.get_eval_dataloaders(
|
|
180
|
+
batch_size=config.sub_batch_size.value(), # type: ignore
|
|
181
|
+
seed=config.seed.value(), # type: ignore
|
|
182
|
+
shuffle=config.shuffle_dataset.value(), # type: ignore
|
|
183
|
+
prefetch_factor=config.row_prefetch.value(), # type: ignore
|
|
184
|
+
num_workers=config.num_row_workers.value(), # type: ignore
|
|
185
|
+
ignore_error=config.ignore_data_error.value(), # type: ignore
|
|
186
|
+
qps=config.data_row_qps.value(),
|
|
187
|
+
instruct_timeout=config.instruct_timeout.value(), # type: ignore
|
|
188
|
+
worker_timeout=config.worker_timeout.value(), # type: ignore
|
|
189
|
+
world_size=get_world_size(),
|
|
190
|
+
rank=get_global_rank(),
|
|
191
|
+
)
|
|
192
|
+
|
|
193
|
+
self.train_task_idxs: list[int] = []
|
|
194
|
+
for idx, loader in enumerate(self.train_data_loaders):
|
|
195
|
+
if loader is not None:
|
|
196
|
+
self.train_task_idxs.append(idx)
|
|
197
|
+
|
|
198
|
+
self.eval_task_idxs: list[int] = []
|
|
199
|
+
for idx, loader in enumerate(self.eval_data_loaders):
|
|
200
|
+
if loader is not None:
|
|
201
|
+
self.eval_task_idxs.append(idx)
|
|
202
|
+
|
|
203
|
+
assert len(self.train_task_idxs) + len(self.eval_task_idxs) > 0, 'At least one train or eval task must be provided.'
|
|
204
|
+
|
|
205
|
+
self.task_weights: list[float] = [w.value() for w in config.task_conf.task_weights] # type: ignore
|
|
206
|
+
self.train_task_weights = [self.task_weights[idx] for idx in self.train_task_idxs]
|
|
207
|
+
self.eval_task_weights = [self.task_weights[idx] for idx in self.eval_task_idxs]
|
|
208
|
+
self.task_rand_g = random.Random(config.seed.value())
|
|
209
|
+
|
|
210
|
+
self.train_iters = [iter(loader) if loader is not None else None for loader in self.train_data_loaders]
|
|
211
|
+
self.eval_iters = [iter(loader) if loader is not None else None for loader in self.eval_data_loaders]
|
|
212
|
+
|
|
213
|
+
def sample_train_task_id(self) -> int:
|
|
214
|
+
'''
|
|
215
|
+
by default, randomly select a task from the task list.
|
|
216
|
+
You can modify this fuction to control the task scheduler.
|
|
217
|
+
Be sure that all ranks are selecting the same task at the same time.
|
|
218
|
+
'''
|
|
219
|
+
return self.task_rand_g.choices(
|
|
220
|
+
list(range(len(self.train_task_idxs))), weights=self.train_task_weights, k=1
|
|
221
|
+
)[0]
|
|
222
|
+
|
|
223
|
+
def sample_eval_task_id(self) -> int:
|
|
224
|
+
'''
|
|
225
|
+
by default, randomly select a task from the task list.
|
|
226
|
+
You can modify this fuction to control the task scheduler.
|
|
227
|
+
Be sure that all ranks are selecting the same task at the same time.
|
|
228
|
+
'''
|
|
229
|
+
return self.task_rand_g.choices(
|
|
230
|
+
list(range(len(self.eval_task_idxs))), weights=self.eval_task_weights, k=1
|
|
231
|
+
)[0]
|
|
232
|
+
|
|
233
|
+
def load_model(self, **kwargs: Any) -> nn.Module:
|
|
234
|
+
return self.model_loader.load_model()
|
|
235
|
+
|
|
236
|
+
def save(self, model: nn.Module, directory: str, step: int, **kwargs: Any) -> None:
|
|
237
|
+
path = Path(directory) / str(step)
|
|
238
|
+
os.makedirs(path, exist_ok=True)
|
|
239
|
+
self.model_loader.save_model(model, path)
|
|
240
|
+
|
|
241
|
+
optimizer_path = path / 'optimizer.pt'
|
|
242
|
+
torch.save(self.optimizer.state_dict(), optimizer_path)
|
|
243
|
+
|
|
244
|
+
scheduler_path = path / 'scheduler.pt'
|
|
245
|
+
torch.save(self.scheduler.state_dict(), scheduler_path)
|
|
246
|
+
|
|
247
|
+
runner_path = path / 'runner.json'
|
|
248
|
+
with open(runner_path, 'w') as f:
|
|
249
|
+
f.write(json.dumps({
|
|
250
|
+
'step': step + 1,
|
|
251
|
+
'epoch': self.epoch,
|
|
252
|
+
'config': self.config.to_dict(),
|
|
253
|
+
}, indent=4))
|
|
254
|
+
|
|
255
|
+
@property
|
|
256
|
+
def model(self) -> DDPWraper:
|
|
257
|
+
model = getattr(self, '_model', None)
|
|
258
|
+
if model is None:
|
|
259
|
+
model = self.load_model()
|
|
260
|
+
if self.config.model_precision == 'fp16':
|
|
261
|
+
model = model.half()
|
|
262
|
+
self._model = DDPWraper(model)
|
|
263
|
+
return self._model
|
|
264
|
+
|
|
265
|
+
def _should_stop(self) -> bool:
|
|
266
|
+
stop_metric = self.config.stop_by.value()
|
|
267
|
+
assert stop_metric in ('step', 'epoch')
|
|
268
|
+
|
|
269
|
+
if stop_metric == 'step':
|
|
270
|
+
total_step = self.config.total_step.value()
|
|
271
|
+
assert total_step is not None
|
|
272
|
+
return self.step > total_step
|
|
273
|
+
elif stop_metric == 'epoch':
|
|
274
|
+
total_epoch = self.config.total_epoch.value()
|
|
275
|
+
assert total_epoch is not None
|
|
276
|
+
return self.epoch >= total_epoch
|
|
277
|
+
else:
|
|
278
|
+
raise ValueError(f'stop_metric must be either "epoch" or "step", got "{stop_metric}" instead.')
|
|
279
|
+
|
|
280
|
+
def _batch_to_device(self, batch: Batch) -> Batch:
|
|
281
|
+
for key, v in batch.items():
|
|
282
|
+
if isinstance(v, Tensor):
|
|
283
|
+
if torch.is_floating_point(v):
|
|
284
|
+
precision = self.config.model_precision.value()
|
|
285
|
+
if precision == 'fp32':
|
|
286
|
+
v = v.to(torch.float32) if v.dtype != torch.float32 else v
|
|
287
|
+
elif precision == 'fp16':
|
|
288
|
+
v = v.to(torch.float16) if v.dtype != torch.float16 else v
|
|
289
|
+
else:
|
|
290
|
+
raise ValueError(f'Unknown model precision "{precision}"')
|
|
291
|
+
assert isinstance(batch, dict)
|
|
292
|
+
batch[key] = v.to(get_default_device()) if v.get_device() != get_default_device() else v
|
|
293
|
+
return batch
|
|
294
|
+
|
|
295
|
+
def _prepare_train(
|
|
296
|
+
self,
|
|
297
|
+
optimizer: Union[Optimizer, None] = None,
|
|
298
|
+
scheduler: Union[LRScheduler, None] = None,
|
|
299
|
+
scaler: Union[GradScaler, None] = None,
|
|
300
|
+
) -> None:
|
|
301
|
+
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
|
|
302
|
+
if optimizer is None:
|
|
303
|
+
self.optimizer = self.config.optimizer.init_optimzier(
|
|
304
|
+
self.model.parameters()
|
|
305
|
+
)
|
|
306
|
+
else:
|
|
307
|
+
self.optimizer = optimizer
|
|
308
|
+
if hasattr(self, '_optimizer_states'):
|
|
309
|
+
self.optimizer.load_state_dict(self._optimizer_states)
|
|
310
|
+
for state in self.optimizer.state.values():
|
|
311
|
+
for k, v in state.items():
|
|
312
|
+
if torch.is_tensor(v):
|
|
313
|
+
state[k] = v.to(get_default_device(), non_blocking=True)
|
|
314
|
+
|
|
315
|
+
if scheduler is None:
|
|
316
|
+
self.scheduler = self.config.lr_scheduler.init_lr_scheduler(
|
|
317
|
+
self.optimizer
|
|
318
|
+
)
|
|
319
|
+
else:
|
|
320
|
+
self.scheduler = scheduler
|
|
321
|
+
if hasattr(self, '_scheduler_states'):
|
|
322
|
+
self.scheduler.load_state_dict(self._scheduler_states)
|
|
323
|
+
|
|
324
|
+
if scaler is None:
|
|
325
|
+
if self.config.use_amp and get_world_size() > 1:
|
|
326
|
+
amp_selection = self.config.amp_precision.value()
|
|
327
|
+
if amp_selection == 'fp16':
|
|
328
|
+
enable_scaler = True
|
|
329
|
+
elif amp_selection == 'bf16':
|
|
330
|
+
enable_scaler = False
|
|
331
|
+
else:
|
|
332
|
+
raise ValueError(f'Unknown amp precision "{amp_selection}"')
|
|
333
|
+
if OLD_GRADSCALER:
|
|
334
|
+
self.scaler = GradScaler(enabled=enable_scaler)
|
|
335
|
+
else:
|
|
336
|
+
self.scaler = GradScaler(device=get_default_device_type(), enabled=enable_scaler) # type: ignore
|
|
337
|
+
else:
|
|
338
|
+
self.scaler = None
|
|
339
|
+
else:
|
|
340
|
+
self.scaler = scaler
|
|
341
|
+
|
|
342
|
+
def _next_train_batch(self, task_idx: int) -> Batch:
|
|
343
|
+
if self.train_data_loaders[task_idx] is None:
|
|
344
|
+
raise ValueError(f'No train dataloader for task {task_idx}')
|
|
345
|
+
|
|
346
|
+
it = self.train_iters[task_idx]
|
|
347
|
+
assert it is not None
|
|
348
|
+
try:
|
|
349
|
+
return next(it)
|
|
350
|
+
except StopIteration:
|
|
351
|
+
self.train_iters[task_idx] = iter(self.train_data_loaders[task_idx]) # type: ignore
|
|
352
|
+
return next(it)
|
|
353
|
+
|
|
354
|
+
@property
|
|
355
|
+
def _wandb(self) -> Run:
|
|
356
|
+
if getattr(self, '_run', None) is None:
|
|
357
|
+
wandb.init(
|
|
358
|
+
project=self.config.project_name.value(),
|
|
359
|
+
name=self.config.run_name.value(),
|
|
360
|
+
config=self.config.to_dict()
|
|
361
|
+
) if get_global_rank() == 0 else ...
|
|
362
|
+
self._run = True
|
|
363
|
+
return self._run # type: ignore
|
|
364
|
+
|
|
365
|
+
@property
|
|
366
|
+
def logger(self) -> Logger:
|
|
367
|
+
if getattr(self, '_logger', None) is None:
|
|
368
|
+
self._logger = getLogger(self.__class__.__name__)
|
|
369
|
+
return self._logger
|
|
370
|
+
|
|
371
|
+
@property
|
|
372
|
+
def epoch(self) -> int:
|
|
373
|
+
epochs = [loader.epoch for loader in self.train_data_loaders if loader is not None]
|
|
374
|
+
return max(epochs) + self.pre_epoch
|
|
375
|
+
|
|
376
|
+
def step_log(self, data: dict[str, Any]) -> None:
|
|
377
|
+
self._wandb
|
|
378
|
+
if get_global_rank() != 0:
|
|
379
|
+
return
|
|
380
|
+
self.logger.critical(f'step:{self.step}\t{data}')
|
|
381
|
+
wandb.log(data, step=self.step)
|
|
382
|
+
|
|
383
|
+
def _one_train_step(self, **kwargs: Any) -> None:
|
|
384
|
+
# clean the gradients
|
|
385
|
+
self.optimizer.zero_grad()
|
|
386
|
+
|
|
387
|
+
# select a task to run in this step
|
|
388
|
+
task_id = self.sample_train_task_id()
|
|
389
|
+
task = self.tasks[task_id]
|
|
390
|
+
|
|
391
|
+
# calculate loss
|
|
392
|
+
running_loss: float = 0.0
|
|
393
|
+
batch_datas: defaultdict[Hashable, List[float]] = defaultdict(list)
|
|
394
|
+
for acc_idx in range(self.config._num_acc_steps):
|
|
395
|
+
with ExitStack() as stack:
|
|
396
|
+
# check whether to use amp
|
|
397
|
+
if all([
|
|
398
|
+
self.config.use_amp,
|
|
399
|
+
self.config.model_precision == 'fp32',
|
|
400
|
+
get_world_size() > 1
|
|
401
|
+
]):
|
|
402
|
+
amp_selection = self.config.amp_precision.value()
|
|
403
|
+
assert amp_selection is not None
|
|
404
|
+
stack.enter_context(autocast(
|
|
405
|
+
device_type='cuda',
|
|
406
|
+
dtype={
|
|
407
|
+
'fp16': torch.float16,
|
|
408
|
+
'bf16': torch.bfloat16
|
|
409
|
+
}[amp_selection],
|
|
410
|
+
))
|
|
411
|
+
|
|
412
|
+
# compute loss for each sub_batch
|
|
413
|
+
subbatch_result = task.train_step(
|
|
414
|
+
model=self.model,
|
|
415
|
+
batch=self._batch_to_device(self._next_train_batch(task_id)),
|
|
416
|
+
step=self.step,
|
|
417
|
+
device=get_local_rank(),
|
|
418
|
+
acc_step=acc_idx,
|
|
419
|
+
)
|
|
420
|
+
if isinstance(subbatch_result, torch.Tensor):
|
|
421
|
+
subbatch_result = {'loss': subbatch_result}
|
|
422
|
+
if 'loss' not in subbatch_result:
|
|
423
|
+
raise KeyError('no loss returned from the batch')
|
|
424
|
+
assert isinstance(subbatch_result, dict)
|
|
425
|
+
loss = subbatch_result['loss']
|
|
426
|
+
assert isinstance(loss, Tensor)
|
|
427
|
+
loss = loss / self.config._num_acc_steps
|
|
428
|
+
running_loss += loss.item()
|
|
429
|
+
for k, v in subbatch_result.items():
|
|
430
|
+
if k == 'loss':
|
|
431
|
+
continue
|
|
432
|
+
if isinstance(v, (float, int)):
|
|
433
|
+
batch_datas[k].append(float(v))
|
|
434
|
+
elif isinstance(v, (list)) and len(v) > 0 and isinstance(v[0], (float, int)):
|
|
435
|
+
batch_datas[k].extend([float(i) for i in v])
|
|
436
|
+
|
|
437
|
+
if self.scaler is not None:
|
|
438
|
+
self.scaler.scale(loss).backward()
|
|
439
|
+
else:
|
|
440
|
+
loss.backward()
|
|
441
|
+
running_loss = dist_avg(running_loss)
|
|
442
|
+
self.step_log({f'{task.__class__.__name__}/train/loss': running_loss})
|
|
443
|
+
self.step_log({'train/epoch': self.epoch})
|
|
444
|
+
for k, v in batch_datas.items():
|
|
445
|
+
try:
|
|
446
|
+
avg = sum(v) / len(v)
|
|
447
|
+
except:
|
|
448
|
+
continue
|
|
449
|
+
self.step_log({f'{task.__class__.__name__}/train/{k}': avg})
|
|
450
|
+
self._pbar_train.set_description(
|
|
451
|
+
f'train loss: {running_loss:.3g}', refresh=True
|
|
452
|
+
)
|
|
453
|
+
|
|
454
|
+
grad_norm_clip_val = self.config.grad_norm_clip.value()
|
|
455
|
+
if grad_norm_clip_val is not None:
|
|
456
|
+
norm = clip_grad_norm_(
|
|
457
|
+
parameters=self.model.parameters(),
|
|
458
|
+
max_norm=grad_norm_clip_val
|
|
459
|
+
).item()
|
|
460
|
+
norm = dist_avg(norm)
|
|
461
|
+
self.step_log({f'{task.__class__.__name__}/train/grad_norm': norm})
|
|
462
|
+
|
|
463
|
+
num_hot_params = 0
|
|
464
|
+
num_freeze_params = 0
|
|
465
|
+
for param in self.model.parameters():
|
|
466
|
+
if param.requires_grad == False:
|
|
467
|
+
param.grad = None
|
|
468
|
+
num_freeze_params += param.numel()
|
|
469
|
+
else:
|
|
470
|
+
num_hot_params += param.numel()
|
|
471
|
+
num_total_params = num_hot_params + num_freeze_params
|
|
472
|
+
if num_total_params == 0:
|
|
473
|
+
raise RuntimeError('The model contains no parameters.')
|
|
474
|
+
|
|
475
|
+
self.step_log({
|
|
476
|
+
f'{task.__class__.__name__}/train/num_hot_params': num_hot_params,
|
|
477
|
+
f'{task.__class__.__name__}/train/num_freeze_params': num_freeze_params,
|
|
478
|
+
f'{task.__class__.__name__}/train/num_total_params': num_total_params,
|
|
479
|
+
f'{task.__class__.__name__}/train/hot_ratio': num_hot_params / num_total_params,
|
|
480
|
+
})
|
|
481
|
+
|
|
482
|
+
if self.scaler is not None:
|
|
483
|
+
self.scaler.step(self.optimizer)
|
|
484
|
+
self.scaler.update()
|
|
485
|
+
else:
|
|
486
|
+
self.optimizer.step()
|
|
487
|
+
current_lr = self.scheduler.get_lr()
|
|
488
|
+
if isinstance(current_lr, list):
|
|
489
|
+
current_lr = current_lr[0]
|
|
490
|
+
self.step_log(
|
|
491
|
+
{f'{task.__class__.__name__}/train/learning_rate': current_lr}
|
|
492
|
+
)
|
|
493
|
+
self.scheduler.step()
|
|
494
|
+
torch.distributed.barrier() if get_world_size() > 1 else ...
|
|
495
|
+
|
|
496
|
+
save_step_freq = self.config.save_step_freq.value()
|
|
497
|
+
assert save_step_freq is not None
|
|
498
|
+
if all(
|
|
499
|
+
[self.step % save_step_freq == 0, get_global_rank() == 0]
|
|
500
|
+
):
|
|
501
|
+
logger.info('begin to save the model')
|
|
502
|
+
self._pbar_train.set_description('begin to save the model', True)
|
|
503
|
+
to_save = self.model.model.module
|
|
504
|
+
ckpt_path = self.config.checkpoint_directory.value()
|
|
505
|
+
assert ckpt_path is not None
|
|
506
|
+
self.save(
|
|
507
|
+
to_save, # type: ignore
|
|
508
|
+
ckpt_path,
|
|
509
|
+
self.step
|
|
510
|
+
)
|
|
511
|
+
logger.info('model saved!')
|
|
512
|
+
self._pbar_train.set_description('model saved!', True)
|
|
513
|
+
torch.distributed.barrier() if get_world_size() > 1 else ...
|
|
514
|
+
|
|
515
|
+
eval_step_freq = self.config.eval_step_freq.value()
|
|
516
|
+
assert eval_step_freq is not None
|
|
517
|
+
if self.step % eval_step_freq == 0:
|
|
518
|
+
logger.info('begin to evaluate the model')
|
|
519
|
+
self.eval()
|
|
520
|
+
logger.info('model evaluated!')
|
|
521
|
+
self.step += 1
|
|
522
|
+
|
|
523
|
+
def train(self, *args: Any, **kwargs: Any) -> None:
|
|
524
|
+
self._prepare_train()
|
|
525
|
+
|
|
526
|
+
self._pbar_train = tqdm(
|
|
527
|
+
total=self.config.total_step.value(),
|
|
528
|
+
position=0,
|
|
529
|
+
dynamic_ncols=True,
|
|
530
|
+
unit='step',
|
|
531
|
+
disable=True if get_global_rank() != 0 else False
|
|
532
|
+
)
|
|
533
|
+
|
|
534
|
+
while not self._should_stop():
|
|
535
|
+
self._one_train_step()
|
|
536
|
+
self._pbar_train.update(1)
|
|
537
|
+
|
|
538
|
+
self._pbar_train.close()
|
|
539
|
+
|
|
540
|
+
def _eval_one_task(self, task_id: int, **kwargs: Any) -> None:
|
|
541
|
+
os.environ['TOKENIZERS_PARALLELISM'] = 'false'
|
|
542
|
+
|
|
543
|
+
dataloader = self.eval_data_loaders[task_id]
|
|
544
|
+
task = self.tasks[task_id]
|
|
545
|
+
|
|
546
|
+
batch_list: list[dict[str, list[Any]]] = []
|
|
547
|
+
with ExitStack() as stack:
|
|
548
|
+
stack.enter_context(torch.no_grad())
|
|
549
|
+
if get_world_size() > 1:
|
|
550
|
+
stack.enter_context(self.model.model.no_sync())
|
|
551
|
+
for batch in tqdm(
|
|
552
|
+
dataloader,
|
|
553
|
+
dynamic_ncols=True,
|
|
554
|
+
desc=f'evaluating task {task.__class__.__name__}',
|
|
555
|
+
position=2,
|
|
556
|
+
disable=True if get_global_rank() != 0 else False,
|
|
557
|
+
leave=False,
|
|
558
|
+
unit='batch'
|
|
559
|
+
):
|
|
560
|
+
result = task.eval_step(
|
|
561
|
+
self.model, # type: ignore
|
|
562
|
+
self._batch_to_device(batch),
|
|
563
|
+
self.step,
|
|
564
|
+
get_default_device()
|
|
565
|
+
)
|
|
566
|
+
batch_list.append(result)
|
|
567
|
+
|
|
568
|
+
if len(batch_list) == 0:
|
|
569
|
+
raise RuntimeError(f'No eval data found in rank {get_global_rank()} '
|
|
570
|
+
f'for task {task.__class__.__name__}')
|
|
571
|
+
cat_result = batch_list[0]
|
|
572
|
+
for batch in batch_list[1:]:
|
|
573
|
+
for k, v in batch.items():
|
|
574
|
+
cat_result[k] += v
|
|
575
|
+
all_result = batch_all_gather(cat_result)
|
|
576
|
+
|
|
577
|
+
with torch.no_grad():
|
|
578
|
+
metrics = task.cal_dev_metric(all_result)
|
|
579
|
+
for k, v in metrics.items():
|
|
580
|
+
self.step_log({f'{task.__class__.__name__}/dev/{k}': v})
|
|
581
|
+
|
|
582
|
+
def eval(self, *args: Any, **kwargs: Any) -> None:
|
|
583
|
+
for task_id in tqdm(
|
|
584
|
+
self.eval_task_idxs,
|
|
585
|
+
dynamic_ncols=True,
|
|
586
|
+
desc=f'evaluating {len(self.tasks)} tasks...',
|
|
587
|
+
position=2,
|
|
588
|
+
disable=True if get_global_rank() != 0 else False,
|
|
589
|
+
leave=False,
|
|
590
|
+
unit='task'
|
|
591
|
+
):
|
|
592
|
+
self._eval_one_task(task_id)
|
|
593
|
+
|
|
594
|
+
def produce(self, *args: Any, **kwargs: Any) -> None:
|
|
595
|
+
raise NotImplementedError('produce method not implemented')
|
|
596
|
+
|
|
597
|
+
def load(self, directory: str | os.PathLike, *args, **kwargs) -> None:
|
|
598
|
+
self.model_loader.model_path = directory
|
|
599
|
+
|
|
600
|
+
optimizer_path = Path(directory) / 'optimizer.pt'
|
|
601
|
+
if optimizer_path.exists():
|
|
602
|
+
self._optimizer_states = torch.load(optimizer_path)
|
|
603
|
+
else:
|
|
604
|
+
logger.warning(f'optimizer.pt not found in {directory}, skip loading optimizer')
|
|
605
|
+
|
|
606
|
+
scheduler_path = Path(directory) / 'scheduler.pt'
|
|
607
|
+
if scheduler_path.exists():
|
|
608
|
+
self._scheduler_states = torch.load(scheduler_path)
|
|
609
|
+
else:
|
|
610
|
+
logger.warning(f'scheduler.pt not found in {directory}, skip loading scheduler')
|
|
611
|
+
|
|
612
|
+
runner_path = Path(directory) / 'runner.json'
|
|
613
|
+
if runner_path.exists():
|
|
614
|
+
with open(runner_path, 'r') as f:
|
|
615
|
+
runner_state = json.load(f)
|
|
616
|
+
self.step = int(runner_state['step'])
|
|
617
|
+
self.pre_epoch = int(runner_state['epoch'])
|
|
618
|
+
else:
|
|
619
|
+
logger.warning(f'runner.json not found in {directory}, skip loading runner state')
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
from typing import Any, Generic, TypeVar
|
|
2
|
+
from abc import ABC, abstractmethod
|
|
3
|
+
import os
|
|
4
|
+
|
|
5
|
+
from hyperargs import Conf, StrArg
|
|
6
|
+
from lmfuser_data.interfaces import SubclassTracer
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class RunerConf(Conf, SubclassTracer):
|
|
10
|
+
|
|
11
|
+
project_name = StrArg('please set a project name')
|
|
12
|
+
run_name = StrArg('please set the name of this run')
|
|
13
|
+
|
|
14
|
+
ConfType = TypeVar('ConfType', bound=RunerConf)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Runner(ABC, Generic[ConfType]):
|
|
18
|
+
def __init__(self, config: ConfType, *args, **kwargs) -> None:
|
|
19
|
+
super().__init__()
|
|
20
|
+
self.config = config
|
|
21
|
+
|
|
22
|
+
@abstractmethod
|
|
23
|
+
def train(self, *args: Any, **kwargs: Any) -> None:
|
|
24
|
+
raise NotImplementedError('train method not implemented')
|
|
25
|
+
|
|
26
|
+
@abstractmethod
|
|
27
|
+
def eval(self, *args: Any, **kwargs: Any) -> None:
|
|
28
|
+
raise NotImplementedError('train method not implemented')
|
|
29
|
+
|
|
30
|
+
@abstractmethod
|
|
31
|
+
def produce(self, *args: Any, **kwargs: Any) -> None:
|
|
32
|
+
raise NotImplementedError('produce method not implemented')
|
|
33
|
+
|
|
34
|
+
@abstractmethod
|
|
35
|
+
def save(self, directory: str | os.PathLike, *args, **kwargs) -> None:
|
|
36
|
+
raise NotImplementedError('produce method not implemented')
|
|
37
|
+
|
|
38
|
+
@abstractmethod
|
|
39
|
+
def load(self, directory: str | os.PathLike, *args, **kwargs) -> None:
|
|
40
|
+
raise NotImplementedError('load method not implemented')
|