Improve: style init validation check

This commit is contained in:
litagin02
2024-01-03 12:50:43 +09:00
parent 33d18e825a
commit 9ee632bb96

View File

@@ -35,13 +35,23 @@ class Model:
self.hps: utils.HParams = utils.get_hparams_from_file(self.config_path) self.hps: utils.HParams = utils.get_hparams_from_file(self.config_path)
self.spk2id: Dict[str, int] = self.hps.data.spk2id self.spk2id: Dict[str, int] = self.hps.data.spk2id
self.id2spk: Dict[int, str] = {v: k for k, v in self.spk2id.items()} self.id2spk: Dict[int, str] = {v: k for k, v in self.spk2id.items()}
self.num_styles: int = self.hps.data.num_styles self.num_styles: int = self.hps.data.num_styles
if hasattr(self.hps.data, "style2id"): if hasattr(self.hps.data, "style2id"):
self.style2id: Dict[str, int] = self.hps.data.style2id self.style2id: Dict[str, int] = self.hps.data.style2id
else: else:
self.style2id: Dict[str, int] = {str(i): i for i in range(self.num_styles)} self.style2id: Dict[str, int] = {str(i): i for i in range(self.num_styles)}
if len(self.style2id) != self.num_styles:
raise ValueError(
f"Number of styles ({self.num_styles}) does not match the number of style2id ({len(self.style2id)})"
)
self.style_vectors: np.ndarray = np.load(self.style_vec_path) self.style_vectors: np.ndarray = np.load(self.style_vec_path)
if self.style_vectors.shape[0] != self.num_styles:
logger.warning(
f"The number of styles ({self.num_styles}) does not match the number of style vectors ({self.style_vectors.shape[0]})"
)
self.net_g: Optional[SynthesizerTrn] = None self.net_g: Optional[SynthesizerTrn] = None
def load_net_g(self): def load_net_g(self):