This commit is contained in:
enhuiz 2023-01-19 00:03:51 +08:00
parent d80ef1d970
commit 2639ba9cab
2 changed files with 8 additions and 13 deletions

View File

@ -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):

View File

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