LMFuser 0.0.7__tar.gz → 0.0.9__tar.gz

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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: LMFuser
3
- Version: 0.0.7
3
+ Version: 0.0.9
4
4
  Summary: The LMFuser training framework.
5
5
  Project-URL: Homepage, https://github.com/TYTTYTTYT/LMFuser
6
6
  Project-URL: Documentation, https://github.com/TYTTYTTYT/LMFuser
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "LMFuser"
7
- version = "0.0.7"
7
+ version = "0.0.9"
8
8
  requires-python = ">= 3.11"
9
9
  description = "The LMFuser training framework."
10
10
  readme = "README.md"
@@ -190,6 +190,20 @@ class DDPRunner(Runner[DDPRunnerConfig]):
190
190
  rank=get_global_rank(),
191
191
  )
192
192
 
193
+ self.test_data_loaders = config.task_conf.get_test_dataloaders(
194
+ batch_size=config.sub_batch_size.value(), # type: ignore
195
+ seed=config.seed.value(), # type: ignore
196
+ shuffle=config.shuffle_dataset.value(), # type: ignore
197
+ prefetch_factor=config.row_prefetch.value(), # type: ignore
198
+ num_workers=config.num_row_workers.value(), # type: ignore
199
+ ignore_error=config.ignore_data_error.value(), # type: ignore
200
+ qps=config.data_row_qps.value(),
201
+ instruct_timeout=config.instruct_timeout.value(), # type: ignore
202
+ worker_timeout=config.worker_timeout.value(), # type: ignore
203
+ world_size=get_world_size(),
204
+ rank=get_global_rank(),
205
+ )
206
+
193
207
  self.train_task_idxs: list[int] = []
194
208
  for idx, loader in enumerate(self.train_data_loaders):
195
209
  if loader is not None:
@@ -200,6 +214,11 @@ class DDPRunner(Runner[DDPRunnerConfig]):
200
214
  if loader is not None:
201
215
  self.eval_task_idxs.append(idx)
202
216
 
217
+ self.test_task_idxs: list[int] = []
218
+ for idx, loader in enumerate(self.test_data_loaders):
219
+ if loader is not None:
220
+ self.test_task_idxs.append(idx)
221
+
203
222
  assert len(self.train_task_idxs) + len(self.eval_task_idxs) > 0, 'At least one train or eval task must be provided.'
204
223
 
205
224
  self.task_weights: list[float] = [w.value() for w in config.task_conf.task_weights] # type: ignore
@@ -210,6 +229,9 @@ class DDPRunner(Runner[DDPRunnerConfig]):
210
229
  self.train_iters = [iter(loader) if loader is not None else None for loader in self.train_data_loaders]
211
230
  self.eval_iters = [iter(loader) if loader is not None else None for loader in self.eval_data_loaders]
212
231
 
232
+ self._all_eval_results: dict[str, list[dict[str, Any]]] = {}
233
+ self._test_results: dict[str, dict[str, Any]] = {}
234
+
213
235
  def sample_train_task_id(self) -> int:
214
236
  '''
215
237
  by default, randomly select a task from the task list.
@@ -261,6 +283,7 @@ class DDPRunner(Runner[DDPRunnerConfig]):
261
283
  logger.critical(f'casting model to fp16')
262
284
  model = model.half()
263
285
  self._model = DDPWraper(model)
286
+ assert self._model is not None
264
287
  return self._model
265
288
 
266
289
  def _should_stop(self) -> bool:
@@ -550,7 +573,7 @@ class DDPRunner(Runner[DDPRunnerConfig]):
550
573
  dataloader,
551
574
  dynamic_ncols=True,
552
575
  desc=f'evaluating task {task.__class__.__name__}',
553
- position=2,
576
+ position=3,
554
577
  disable=True if get_global_rank() != 0 else False,
555
578
  leave=False,
556
579
  unit='batch'
@@ -577,6 +600,64 @@ class DDPRunner(Runner[DDPRunnerConfig]):
577
600
  for k, v in metrics.items():
578
601
  self.step_log({f'{task.__class__.__name__}/dev/{k}': v})
579
602
 
603
+ if task.__class__.__name__ not in self._all_eval_results:
604
+ self._all_eval_results[task.__class__.__name__] = []
605
+ self._all_eval_results[task.__class__.__name__].append({'step': self.step, 'metrics': metrics})
606
+
607
+ def _test_one_task(self, task_id: int, **kwargs: Any) -> None:
608
+ os.environ['TOKENIZERS_PARALLELISM'] = 'false'
609
+
610
+ dataloader = self.test_data_loaders[task_id]
611
+ task = self.tasks[task_id]
612
+
613
+ ckpt_path = self.config.checkpoint_directory.value()
614
+ assert ckpt_path is not None
615
+ test_model_path = task.set_test_model_path(ckpt_path, self._all_eval_results)
616
+ if test_model_path is not None:
617
+ if isinstance(test_model_path, int):
618
+ test_model_path = Path(ckpt_path) / str(test_model_path)
619
+
620
+ self.model_loader.model_path = test_model_path
621
+ self._model = None
622
+
623
+ batch_list: list[dict[str, list[Any]]] = []
624
+ with ExitStack() as stack:
625
+ stack.enter_context(torch.no_grad())
626
+ if get_world_size() > 1:
627
+ stack.enter_context(self.model.model.no_sync())
628
+ for batch in tqdm(
629
+ dataloader,
630
+ dynamic_ncols=True,
631
+ desc=f'testing task {task.__class__.__name__}',
632
+ position=4,
633
+ disable=True if get_global_rank() != 0 else False,
634
+ leave=False,
635
+ unit='batch'
636
+ ):
637
+ result = task.eval_step(
638
+ self.model, # type: ignore
639
+ self._batch_to_device(batch),
640
+ self.step,
641
+ get_default_device()
642
+ )
643
+ batch_list.append(result)
644
+
645
+ if len(batch_list) == 0:
646
+ raise RuntimeError(f'No test data found in rank {get_global_rank()} '
647
+ f'for task {task.__class__.__name__}')
648
+ cat_result = batch_list[0]
649
+ for batch in batch_list[1:]:
650
+ for k, v in batch.items():
651
+ cat_result[k] += v
652
+ all_result = batch_all_gather(cat_result)
653
+
654
+ with torch.no_grad():
655
+ metrics = task.cal_dev_metric(all_result)
656
+ for k, v in metrics.items():
657
+ self.step_log({f'{task.__class__.__name__}/test/{k}': v})
658
+
659
+ self._test_results[task.__class__.__name__] = {'step': self.step, 'metrics': metrics}
660
+
580
661
  def eval(self, *args: Any, **kwargs: Any) -> None:
581
662
  for task_id in tqdm(
582
663
  self.eval_task_idxs,
@@ -589,6 +670,48 @@ class DDPRunner(Runner[DDPRunnerConfig]):
589
670
  ):
590
671
  self._eval_one_task(task_id)
591
672
 
673
+ if get_global_rank() != 0:
674
+ return
675
+ ckpt_path = self.config.checkpoint_directory.value()
676
+ assert ckpt_path is not None
677
+ try:
678
+ eval_results = json.dumps(self._all_eval_results, indent=2)
679
+ ckpt_dir = Path(ckpt_path)
680
+ if not ckpt_dir.exists():
681
+ ckpt_dir.mkdir(parents=True)
682
+ eval_path = ckpt_dir / 'eval_results.json'
683
+ with open(eval_path, 'w') as f:
684
+ f.write(eval_results)
685
+ except:
686
+ logger.warning(f'Failed to save eval results to {ckpt_path}')
687
+
688
+ def test(self, *args: Any, **kwargs: Any) -> None:
689
+ for task_id in tqdm(
690
+ self.test_task_idxs,
691
+ dynamic_ncols=True,
692
+ desc=f'testing {len(self.test_task_idxs)} tasks...',
693
+ position=3,
694
+ disable=True if get_global_rank() != 0 else False,
695
+ leave=False,
696
+ unit='task'
697
+ ):
698
+ self._test_one_task(task_id)
699
+
700
+ if get_global_rank() != 0:
701
+ return
702
+ ckpt_path = self.config.checkpoint_directory.value()
703
+ assert ckpt_path is not None
704
+ try:
705
+ test_results = json.dumps(self._test_results, indent=2)
706
+ ckpt_dir = Path(ckpt_path)
707
+ if not ckpt_dir.exists():
708
+ ckpt_dir.mkdir(parents=True)
709
+ test_path = ckpt_dir / 'test_results.json'
710
+ with open(test_path, 'w') as f:
711
+ f.write(test_results)
712
+ except:
713
+ logger.warning(f'Failed to save test results to {ckpt_path}')
714
+
592
715
  def produce(self, *args: Any, **kwargs: Any) -> None:
593
716
  raise NotImplementedError('produce method not implemented')
594
717
 
@@ -25,7 +25,11 @@ class Runner(ABC, Generic[ConfType]):
25
25
 
26
26
  @abstractmethod
27
27
  def eval(self, *args: Any, **kwargs: Any) -> None:
28
- raise NotImplementedError('train method not implemented')
28
+ raise NotImplementedError('eval method not implemented')
29
+
30
+ @abstractmethod
31
+ def test(self, *args: Any, **kwargs: Any) -> None:
32
+ raise NotImplementedError('test method not implemented')
29
33
 
30
34
  @abstractmethod
31
35
  def produce(self, *args: Any, **kwargs: Any) -> None:
@@ -1,5 +1,5 @@
1
1
  from typing import Any, Callable
2
- from collections.abc import Iterable
2
+ from collections.abc import Iterable, Iterator
3
3
 
4
4
  import torch
5
5
  from torch import nn
@@ -12,8 +12,32 @@ from hyperargs import Conf, StrArg, FloatArg, IntArg, OptionArg, add_dependency,
12
12
  def scanner_type_list() -> list[str]:
13
13
  return list(Scanner.all_subclass_names())
14
14
 
15
+
16
+ class EmptyDataLoader:
17
+ '''
18
+ EmptyDataLoader is a class that provides empty data for running with no data requierment.
19
+ '''
20
+ def __init__(self, init_step: int = 0) -> None:
21
+ self.init_step = init_step
22
+
23
+ @property
24
+ def epoch(self) -> int:
25
+ return 0
26
+
27
+ def __iter__(self) -> Iterator[Batch]:
28
+ def it_wrap() -> Iterator[Batch]:
29
+ while True:
30
+ yield {'step': torch.tensor(self.init_step)}
31
+ self.init_step += 1
32
+ return it_wrap()
33
+
34
+
15
35
  @add_dependency('num_train_data_path', 'train_data_path_list')
16
36
  @add_dependency('num_train_data_path', 'train_data_weights')
37
+ @add_dependency('num_eval_data_path', 'eval_data_path_list')
38
+ @add_dependency('num_eval_data_path', 'eval_data_weights')
39
+ @add_dependency('num_test_data_path', 'test_data_path_list')
40
+ @add_dependency('num_test_data_path', 'test_data_weights')
17
41
  class TaskBase(Conf, SubclassTracer):
18
42
  num_train_data_path = IntArg(1, min_value=0)
19
43
  train_data_path_list = [StrArg('Enther the path to the data file.')]
@@ -23,13 +47,19 @@ class TaskBase(Conf, SubclassTracer):
23
47
  eval_data_path_list = [StrArg('Enther the path to the data file.')]
24
48
  eval_data_weights = [FloatArg(1.0, min_value=0.0, max_value=1.0)]
25
49
 
50
+ num_test_data_path = IntArg(0, min_value=0)
51
+ test_data_path_list: list[StrArg] = []
52
+ test_data_weights: list[FloatArg] = []
53
+
26
54
  scanner_type = OptionArg(default='C4Scanner', option_fn=scanner_type_list)
27
55
 
28
- train_dataloader_type = OptionArg(default='single file', options=['single file', 'sharded'])
29
- eval_dataloader_type = OptionArg(default='single file', options=['single file', 'sharded'])
56
+ train_dataloader_type = OptionArg(default='single file', options=['single file', 'sharded', 'empty'])
57
+ eval_dataloader_type = OptionArg(default='single file', options=['single file', 'sharded', 'empty'])
58
+ test_dataloader_type = OptionArg(default='single file', options=['single file', 'sharded', 'empty'])
30
59
 
31
- _train_dataloader: DataLoader | None | PyTorchDataLoader = None
32
- _eval_dataloader: DataLoader | None | PyTorchDataLoader = None
60
+ _train_dataloader: DataLoader | None | PyTorchDataLoader | EmptyDataLoader = None
61
+ _eval_dataloader: DataLoader | None | PyTorchDataLoader | EmptyDataLoader = None
62
+ _test_dataloader: DataLoader | None | PyTorchDataLoader | EmptyDataLoader = None
33
63
 
34
64
  @monitor_on('num_train_data_path')
35
65
  def set_train_path_list(self) -> None:
@@ -50,8 +80,19 @@ class TaskBase(Conf, SubclassTracer):
50
80
  self.eval_data_path_list = self.eval_data_path_list[:num]
51
81
  self.eval_data_weights = self.eval_data_weights[:num]
52
82
  elif len(self.eval_data_path_list) < num:
53
- self.eval_data_path_list += [StrArg('Enther the path to the data file.')] * (num - len(self.train_data_path_list))
54
- self.eval_data_weights += [FloatArg(1.0, min_value=0.0, max_value=1.0)] * (num - len(self.train_data_weights))
83
+ self.eval_data_path_list += [StrArg('Enther the path to the data file.')] * (num - len(self.eval_data_path_list))
84
+ self.eval_data_weights += [FloatArg(1.0, min_value=0.0, max_value=1.0)] * (num - len(self.eval_data_weights))
85
+
86
+ @monitor_on('num_test_data_path')
87
+ def set_test_path_list(self) -> None:
88
+ num = self.num_test_data_path.value()
89
+ assert isinstance(num, int)
90
+ if len(self.test_data_path_list) > num:
91
+ self.test_data_path_list = self.test_data_path_list[:num]
92
+ self.test_data_path_list = self.test_data_path_list[:num]
93
+ elif len(self.test_data_path_list) < num:
94
+ self.test_data_path_list += [StrArg('Enther the path to the data file.')] * (num - len(self.test_data_path_list))
95
+ self.test_data_weights += [FloatArg(1.0, min_value=0.0, max_value=1.0)] * (num - len(self.test_data_weights))
55
96
 
56
97
  def _get_train_dataloader(
57
98
  self,
@@ -66,7 +107,7 @@ class TaskBase(Conf, SubclassTracer):
66
107
  num_workers: int,
67
108
  rank: int,
68
109
  world_size: int
69
- ) -> None | DataLoader | PyTorchDataLoader:
110
+ ) -> None | DataLoader | PyTorchDataLoader | EmptyDataLoader:
70
111
  if self.num_train_data_path.value() == 0:
71
112
  return None
72
113
  if self._train_dataloader is not None:
@@ -113,6 +154,8 @@ class TaskBase(Conf, SubclassTracer):
113
154
  collate_fn=self.get_collate_fn(),
114
155
  drop_last=False
115
156
  )
157
+ elif dataloader_type == 'empty':
158
+ self._train_dataloader = EmptyDataLoader(init_step=0)
116
159
  else:
117
160
  raise ValueError(f'Unknown dataloader type: {dataloader_type}')
118
161
 
@@ -131,7 +174,7 @@ class TaskBase(Conf, SubclassTracer):
131
174
  num_workers: int,
132
175
  rank: int,
133
176
  world_size: int
134
- ) -> None | DataLoader | PyTorchDataLoader:
177
+ ) -> None | DataLoader | PyTorchDataLoader | EmptyDataLoader:
135
178
  if self.num_eval_data_path.value() == 0:
136
179
  return None
137
180
  if self._eval_dataloader is not None:
@@ -181,11 +224,83 @@ class TaskBase(Conf, SubclassTracer):
181
224
  collate_fn=self.get_collate_fn(),
182
225
  drop_last=False
183
226
  )
227
+ elif dataloader_type == 'empty':
228
+ self._eval_dataloader = EmptyDataLoader(init_step=0)
184
229
  else:
185
230
  raise ValueError(f'Unknown dataloader type: {dataloader_type}')
186
231
 
187
232
  return self._eval_dataloader
188
233
 
234
+ def _get_test_dataloader(
235
+ self,
236
+ batch_size: int,
237
+ seed: int,
238
+ shuffle: bool,
239
+ prefetch_factor: int,
240
+ ignore_error: bool,
241
+ qps: float | None,
242
+ instruct_timeout: float,
243
+ worker_timeout: float,
244
+ num_workers: int,
245
+ rank: int,
246
+ world_size: int
247
+ ) -> None | DataLoader | PyTorchDataLoader | EmptyDataLoader:
248
+ if self.num_test_data_path.value() == 0:
249
+ return None
250
+ if self._test_dataloader is not None:
251
+ return self._test_dataloader
252
+ path_list = [p.value() for p in self.test_data_path_list]
253
+ weight_list = [w.value() for w in self.test_data_weights]
254
+ scanner_type = self.scanner_type.value()
255
+ assert scanner_type is not None, 'scanner_type is None'
256
+
257
+ dataloader_type = self.test_dataloader_type.value()
258
+ assert dataloader_type is not None, 'dataloader_type is None'
259
+
260
+ dataloader_type = self.test_dataloader_type.value()
261
+ assert dataloader_type in ('sharded', 'single file'), f'Unknown dataloader type: {dataloader_type}'
262
+
263
+ if dataloader_type == 'sharded':
264
+ self._test_dataloader = DataLoader(
265
+ batch_size=batch_size,
266
+ path_list=path_list, # type: ignore
267
+ distributor_weights=weight_list, # type: ignore
268
+ scanner_type=Scanner.get_subclass(scanner_type),
269
+ seed=seed,
270
+ shuffle=shuffle,
271
+ pre_fetch_factor=prefetch_factor,
272
+ ignore_error=ignore_error,
273
+ qps=qps,
274
+ instruct_timeout=instruct_timeout,
275
+ worker_timeout=worker_timeout,
276
+ num_workers=num_workers,
277
+ map_fn=self.get_row_processor(),
278
+ flow_fn=self.get_flow_processor(),
279
+ batch_map_fn=self.get_batch_processor(),
280
+ rank_idx=rank,
281
+ num_ranks=world_size,
282
+ )
283
+ elif dataloader_type == 'single file':
284
+ self._test_dataloader = PyTorchDataLoader(
285
+ batch_size=batch_size,
286
+ path_list=path_list, # type: ignore
287
+ scanner_type=Scanner.get_subclass(scanner_type),
288
+ seed=seed,
289
+ shuffle=shuffle,
290
+ pre_fetch_factor=prefetch_factor,
291
+ num_workers=num_workers,
292
+ num_ranks=world_size,
293
+ rank_idx=rank,
294
+ collate_fn=self.get_collate_fn(),
295
+ drop_last=False
296
+ )
297
+ elif dataloader_type == 'empty':
298
+ self._test_dataloader = EmptyDataLoader(init_step=0)
299
+ else:
300
+ raise ValueError(f'Unknown dataloader type: {dataloader_type}')
301
+
302
+ return self._test_dataloader
303
+
189
304
  def train_step(
190
305
  self, model: nn.Module,
191
306
  batch: Batch,
@@ -209,6 +324,17 @@ class TaskBase(Conf, SubclassTracer):
209
324
  def cal_dev_metric(self, eval_outputs: dict[str, list[Any]]) -> dict[str, Any]:
210
325
  raise NotImplementedError('Please implement this method in child class')
211
326
 
327
+ def set_test_model_path(self, checkpoints_path: str, dev_metrics: dict[str, list[dict[str, Any]]]) -> str | int | None:
328
+ '''Set the path of the test model based on the dev metrics.
329
+ Args:
330
+ checkpoints_path (str): The path to the checkpoints directory.
331
+ dev_metrics (dict[str, list[dict[str, Any]]]): The development metrics for each task.
332
+
333
+ Returns:
334
+ str | int | None: The path or the checkpoint step to the test model, or None to use the last step model.
335
+ '''
336
+ return None
337
+
212
338
  def get_row_processor(self) -> Callable[[Row], Row] | None:
213
339
  return None
214
340
 
@@ -277,7 +403,7 @@ class Tasks(Conf):
277
403
  num_workers: int,
278
404
  rank: int,
279
405
  world_size: int
280
- ) -> list[DataLoader | None | PyTorchDataLoader]:
406
+ ) -> list[DataLoader | None | PyTorchDataLoader | EmptyDataLoader]:
281
407
  return [
282
408
  task.conf._get_train_dataloader(
283
409
  batch_size=batch_size,
@@ -308,7 +434,7 @@ class Tasks(Conf):
308
434
  num_workers: int,
309
435
  rank: int,
310
436
  world_size: int
311
- ) -> list[DataLoader | None | PyTorchDataLoader]:
437
+ ) -> list[DataLoader | None | PyTorchDataLoader | EmptyDataLoader]:
312
438
  return [
313
439
  task.conf._get_eval_dataloader(
314
440
  batch_size=batch_size,
@@ -326,9 +452,33 @@ class Tasks(Conf):
326
452
  for task in self.tasks
327
453
  ]
328
454
 
329
- if __name__ == '__main__':
330
- class TaskTemp(Task):
331
- pass
332
-
333
- conf = Tasks.parse_command_line()
334
- print(conf)
455
+ def get_test_dataloaders(
456
+ self,
457
+ batch_size: int,
458
+ seed: int,
459
+ shuffle: bool,
460
+ prefetch_factor: int,
461
+ ignore_error: bool,
462
+ qps: float | None,
463
+ instruct_timeout: float,
464
+ worker_timeout: float,
465
+ num_workers: int,
466
+ rank: int,
467
+ world_size: int
468
+ ) -> list[DataLoader | None | PyTorchDataLoader | EmptyDataLoader]:
469
+ return [
470
+ task.conf._get_test_dataloader(
471
+ batch_size=batch_size,
472
+ seed=seed,
473
+ shuffle=shuffle,
474
+ prefetch_factor=prefetch_factor,
475
+ ignore_error=ignore_error,
476
+ qps=qps,
477
+ instruct_timeout=instruct_timeout,
478
+ worker_timeout=worker_timeout,
479
+ num_workers=num_workers,
480
+ rank=rank,
481
+ world_size=world_size,
482
+ )
483
+ for task in self.tasks
484
+ ]
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes