Remove unused test dl

This commit is contained in:
enhuiz 2023-01-18 23:52:49 +08:00
parent 2e9f5030a3
commit d80ef1d970
2 changed files with 17 additions and 34 deletions

View File

@ -244,20 +244,14 @@ def _load_train_val_paths():
return train_paths, val_paths return train_paths, val_paths
def _load_test_paths():
test_paths = []
for data_dir in cfg.test_data_dirs:
test_paths.extend(data_dir.rglob("*.phn.txt"))
test_paths = sorted(test_paths)
return test_paths
@cfg.diskcache() @cfg.diskcache()
def create_datasets(): def create_datasets():
train_paths, val_paths = _load_train_val_paths() train_paths, val_paths = _load_train_val_paths()
test_paths = _load_test_paths()
train_dataset = VALLEDatset(train_paths, training=True) train_dataset = VALLEDatset(
train_paths,
training=True,
)
val_dataset = VALLEDatset( val_dataset = VALLEDatset(
val_paths, val_paths,
@ -269,41 +263,32 @@ def create_datasets():
val_dataset.interleaved_reorder_(_get_spkr_name) val_dataset.interleaved_reorder_(_get_spkr_name)
val_dataset.head_(cfg.max_num_val) val_dataset.head_(cfg.max_num_val)
test_dataset = VALLEDatset( return train_dataset, val_dataset
test_paths,
train_dataset.phone_symmap,
train_dataset.spkr_symmap,
extra_paths_by_spkr_name=train_dataset.paths_by_spkr_name,
)
return train_dataset, val_dataset, test_dataset
def create_train_val_dataloader(): def create_train_val_dataloader():
train_dataset, val_dataset, test_dataset = create_datasets() train_dataset, val_dataset = create_datasets()
train_dl = _create_dl(train_dataset, training=True) train_dl = _create_dl(train_dataset, training=True)
val_dl = _create_dl(val_dataset, training=False) val_dl = _create_dl(val_dataset, training=False)
test_dl = _create_dl(test_dataset, training=False)
_logger.info(str(train_dataset.phone_symmap)) _logger.info(str(train_dataset.phone_symmap))
_logger.info(str(train_dataset.spkr_symmap)) _logger.info(str(train_dataset.spkr_symmap))
_logger.info(f"#samples (train): {len(train_dataset)}.") _logger.info(f"#samples (train): {len(train_dataset)}.")
_logger.info(f"#samples (val): {len(val_dataset)}.") _logger.info(f"#samples (val): {len(val_dataset)}.")
_logger.info(f"#samples (test): {len(test_dataset)}.")
train_for_val_dataset = copy.deepcopy(train_dataset) subtrain_dataset = copy.deepcopy(train_dataset)
train_for_val_dataset.interleaved_reorder_(_get_spkr_name) subtrain_dataset.interleaved_reorder_(_get_spkr_name)
train_for_val_dataset.head_(cfg.max_num_val) subtrain_dataset.head_(cfg.max_num_val)
train_for_val_dataset.training_(False) subtrain_dataset.training_(False)
train_for_val_dl = _create_dl(train_for_val_dataset, training=False) subtrain_dl = _create_dl(subtrain_dataset, training=False)
assert isinstance(train_for_val_dl.dataset, VALLEDatset) assert isinstance(subtrain_dl.dataset, VALLEDatset)
return train_dl, train_for_val_dl, val_dl, test_dl return train_dl, subtrain_dl, val_dl
if __name__ == "__main__": if __name__ == "__main__":
train_dl, train_for_val_dl, val_dl, test_dl = create_train_val_dataloader() train_dl, subtrain_dl, val_dl = create_train_val_dataloader()
sample = train_dl.dataset[0] sample = train_dl.dataset[0]
print(sample) print(sample)

View File

@ -30,7 +30,7 @@ def load_engines():
def main(): def main():
setup_logging(cfg.log_dir) setup_logging(cfg.log_dir)
train_dl, train_for_val_dl, val_dl, test_dl = create_train_val_dataloader() train_dl, subtrain_dl, val_dl = create_train_val_dataloader()
def train_feeder(engines, batch, name): def train_feeder(engines, batch, name):
model = engines["model"] model = engines["model"]
@ -68,8 +68,7 @@ def main():
log_dir = cfg.log_dir / str(engines.global_step) / name log_dir = cfg.log_dir / str(engines.global_step) / name
stats = defaultdict(list) stats = defaultdict(list)
for batch in tqdm(dl): for batch in tqdm(dl):
batch: dict batch: dict = to_device(batch, cfg.device)
batch = to_device(batch, cfg.device)
if cfg.model.startswith("ar"): if cfg.model.startswith("ar"):
resp_list = model( resp_list = model(
@ -114,9 +113,8 @@ def main():
_logger.info(f"{json.dumps(stats)}.") _logger.info(f"{json.dumps(stats)}.")
def eval_fn(engines): def eval_fn(engines):
run_eval(engines, "train_for_val", train_for_val_dl) run_eval(engines, "subtrain", subtrain_dl)
run_eval(engines, "val", val_dl) run_eval(engines, "val", val_dl)
run_eval(engines, "test", test_dl)
trainer.train( trainer.train(
engines_loader=load_engines, engines_loader=load_engines,