Simplify
This commit is contained in:
parent
d80ef1d970
commit
2639ba9cab
@ -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):
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user