From 2639ba9caba7b3dbfdd013069d2aa7900b10e089 Mon Sep 17 00:00:00 2001 From: enhuiz Date: Thu, 19 Jan 2023 00:03:51 +0800 Subject: [PATCH] Simplify --- vall_e/config.py | 1 - vall_e/data.py | 20 ++++++++------------ 2 files changed, 8 insertions(+), 13 deletions(-) diff --git a/vall_e/config.py b/vall_e/config.py index 4692b7f..e967871 100644 --- a/vall_e/config.py +++ b/vall_e/config.py @@ -11,7 +11,6 @@ from .utils import Config as ConfigBase class Config(ConfigBase): data_root: Path = Path("data") data_dirs: list[Path] = field(default_factory=lambda: []) - test_data_dirs: list[Path] = field(default_factory=lambda: []) @property def sample_rate(self): diff --git a/vall_e/data.py b/vall_e/data.py index f855bf9..80d539a 100644 --- a/vall_e/data.py +++ b/vall_e/data.py @@ -71,10 +71,6 @@ def _validate(path, min_phones, max_phones): return True -def _get_spkr_name(path) -> str: - return cfg.get_spkr(path) - - class VALLEDatset(Dataset): def __init__( self, @@ -102,14 +98,14 @@ class VALLEDatset(Dataset): self.paths_by_spkr_name = self._get_paths_by_spkr_name(extra_paths_by_spkr_name) self.paths = [ - p for p in self.paths if len(self.paths_by_spkr_name[_get_spkr_name(p)]) > 1 + p for p in self.paths if len(self.paths_by_spkr_name[cfg.get_spkr(p)]) > 1 ] if len(self.paths) == 0: raise ValueError("No valid path is found. ") if training: - self.sampler = Sampler(self.paths, [_get_spkr_name]) + self.sampler = Sampler(self.paths, [cfg.get_spkr]) else: self.sampler = None @@ -117,7 +113,7 @@ class VALLEDatset(Dataset): ret = defaultdict(list) for path in self.paths: if _get_quant_path(path).exists(): - ret[_get_spkr_name(path)].append(path) + ret[cfg.get_spkr(path)].append(path) for k, v in extra_paths_by_spkr_name.items(): ret[k].extend(v) return {**ret} @@ -132,7 +128,7 @@ class VALLEDatset(Dataset): @cached_property def spkrs(self): - return sorted({_get_spkr_name(path) for path in self.paths}) + return sorted({cfg.get_spkr(path) for path in self.paths}) def _get_spkr_symmap(self): return {s: i for i, s in enumerate(self.spkrs)} @@ -165,7 +161,7 @@ class VALLEDatset(Dataset): else: path = self.paths[index] - spkr_name = _get_spkr_name(path) + spkr_name = cfg.get_spkr(path) text = torch.tensor([*map(self.phone_symmap.get, _get_phones(path))]) proms = self.sample_prompts(spkr_name, ignore=path) resps = _load_quants(path) @@ -228,7 +224,7 @@ def _load_train_val_paths(): if len(paths) == 0: raise RuntimeError(f"Failed to find any .qnt.pt file in {cfg.data_dirs}.") - pairs = sorted([(_get_spkr_name(p), p) for p in paths]) + pairs = sorted([(cfg.get_spkr(p), p) for p in paths]) del paths for _, group in groupby(pairs, lambda pair: pair[0]): @@ -260,7 +256,7 @@ def create_datasets(): extra_paths_by_spkr_name=train_dataset.paths_by_spkr_name, ) - val_dataset.interleaved_reorder_(_get_spkr_name) + val_dataset.interleaved_reorder_(cfg.get_spkr) val_dataset.head_(cfg.max_num_val) return train_dataset, val_dataset @@ -279,7 +275,7 @@ def create_train_val_dataloader(): _logger.info(f"#samples (val): {len(val_dataset)}.") subtrain_dataset = copy.deepcopy(train_dataset) - subtrain_dataset.interleaved_reorder_(_get_spkr_name) + subtrain_dataset.interleaved_reorder_(cfg.get_spkr) subtrain_dataset.head_(cfg.max_num_val) subtrain_dataset.training_(False) subtrain_dl = _create_dl(subtrain_dataset, training=False)