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.
- {lmfuser-0.0.7 → lmfuser-0.0.9}/PKG-INFO +1 -1
- {lmfuser-0.0.7 → lmfuser-0.0.9}/pyproject.toml +1 -1
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/runners/ddp_runner.py +124 -1
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/runners/runner.py +5 -1
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/task.py +167 -17
- {lmfuser-0.0.7 → lmfuser-0.0.9}/.github/workflows/python-publish.yml +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/.gitignore +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/.vscode/settings.json +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/LICENSE +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/README.md +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/__init__.py +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/model_loader.py +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/optimizers.py +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/runners/__init__.py +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/schedulers.py +0 -0
- {lmfuser-0.0.7 → lmfuser-0.0.9}/src/lmfuser/utils.py +0 -0
|
@@ -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=
|
|
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('
|
|
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.
|
|
54
|
-
self.eval_data_weights += [FloatArg(1.0, min_value=0.0, max_value=1.0)] * (num - len(self.
|
|
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
|
-
|
|
330
|
-
|
|
331
|
-
|
|
332
|
-
|
|
333
|
-
|
|
334
|
-
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|