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.
@@ -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')