| 152 | return batch |
| 153 | |
| 154 | def load_test_inputs(self, test_input_dir, spk_id=0): |
| 155 | inp_wav_paths = glob.glob(f'{test_input_dir}/*.wav') + glob.glob(f'{test_input_dir}/*.mp3') |
| 156 | sizes = [] |
| 157 | items = [] |
| 158 | |
| 159 | binarizer_cls = hparams.get("binarizer_cls", 'data_gen.tts.base_binarizerr.BaseBinarizer') |
| 160 | pkg = ".".join(binarizer_cls.split(".")[:-1]) |
| 161 | cls_name = binarizer_cls.split(".")[-1] |
| 162 | binarizer_cls = getattr(importlib.import_module(pkg), cls_name) |
| 163 | binarization_args = hparams['binarization_args'] |
| 164 | |
| 165 | for wav_fn in inp_wav_paths: |
| 166 | item_name = os.path.basename(wav_fn) |
| 167 | ph = txt = tg_fn = '' |
| 168 | wav_fn = wav_fn |
| 169 | encoder = None |
| 170 | item = binarizer_cls.process_item(item_name, ph, txt, tg_fn, wav_fn, spk_id, encoder, binarization_args) |
| 171 | items.append(item) |
| 172 | sizes.append(item['len']) |
| 173 | return items, sizes |