Skip to content

Commit d363a8b

Browse files
committed
feat(tts): resolve evalset paths from root
Add --root to MagpieTTS inference so evalset manifest_path and audio_dir entries can remain relative. Also use the evalset language field for evaluation, preserving whisper_language as a legacy fallback, which is required for multilingual Parakeet target_lang selection. Signed-off-by: quanpham <youngkwan199@gmail.com>
1 parent ed74713 commit d363a8b

1 file changed

Lines changed: 20 additions & 20 deletions

File tree

examples/tts/magpietts_inference.py

Lines changed: 20 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -153,10 +153,7 @@ def create_formatted_metrics_mean_ci(metrics_mean_ci: dict) -> dict:
153153
return metrics_mean_ci
154154

155155

156-
def filter_datasets(
157-
dataset_meta_info: dict,
158-
datasets: Optional[List[str]],
159-
) -> List[str]:
156+
def filter_datasets(dataset_meta_info: dict, datasets: Optional[List[str]]) -> List[str]:
160157
"""Select datasets from the dataset meta info."""
161158
if datasets is None:
162159
# Dataset filtering not specified, return all datasets.
@@ -244,11 +241,13 @@ def run_inference_and_evaluation(
244241

245242
meta = dataset_meta_info[dataset]
246243
manifest_records = read_manifest(meta['manifest_path'])
247-
language = meta.get('whisper_language', 'en')
244+
# `language` drives all language-specific eval logic (ASR target_lang, Whisper prompt,
245+
# text normalization). `whisper_language` is kept as a fallback for legacy evalsets.
246+
language = meta.get('language', meta.get('whisper_language', 'en'))
248247

249248
# Prepare dataset metadata (remove evaluation-specific keys)
250249
dataset_meta_for_dl = copy.deepcopy(meta)
251-
for key in ["whisper_language", "load_cached_codes_if_available"]:
250+
for key in ["language", "whisper_language", "load_cached_codes_if_available"]:
252251
dataset_meta_for_dl.pop(key, None)
253252

254253
# Setup output directories
@@ -495,10 +494,14 @@ def _add_common_args(parser: argparse.ArgumentParser) -> None:
495494
help='Path to dataset configuration JSON file',
496495
)
497496
data_group.add_argument(
498-
'--datasets_base_path',
499-
type=Path,
500-
default=None,
501-
help='Optional base path that paths in the "datasets_json_path" file are relative to',
497+
'--root',
498+
type=str,
499+
default='',
500+
help=(
501+
'Root directory the evalset relative paths are resolved against. '
502+
'Each entry\'s "manifest_path" and "audio_dir" are joined onto this root. '
503+
'Defaults to empty (treat paths as absolute / cwd-relative).'
504+
),
502505
)
503506
data_group.add_argument(
504507
'--datasets',
@@ -604,11 +607,6 @@ def _add_easy_magpie_args(parser: argparse.ArgumentParser) -> None:
604607
default=None,
605608
help='Override path to the phoneme tokenizer file (overrides the path stored in the checkpoint config)',
606609
)
607-
group.add_argument(
608-
'--disable_cas_for_context_text',
609-
action='store_true',
610-
help='Skip CAS embeddings for context text when loading legacy EasyMagpieTTS models',
611-
)
612610

613611

614612
def create_argument_parser() -> argparse.ArgumentParser:
@@ -670,9 +668,13 @@ def main(argv=None):
670668
if args.deterministic:
671669
seed_all(seed=9)
672670

673-
dataset_meta_info = load_evalset_config(
674-
config_path=args.datasets_json_path, dataset_base_path=args.datasets_base_path
675-
)
671+
dataset_meta_info = load_evalset_config(args.datasets_json_path)
672+
# Resolve relative evalset paths against --root so the checked-in config stays portable.
673+
if args.root:
674+
for _meta in dataset_meta_info.values():
675+
for _key in ("manifest_path", "audio_dir"):
676+
if _meta.get(_key):
677+
_meta[_key] = os.path.join(args.root, _meta[_key])
676678
datasets = filter_datasets(dataset_meta_info, args.datasets)
677679
logging.info(f"Loaded {len(datasets)} datasets: {', '.join(datasets)}")
678680

@@ -724,7 +726,6 @@ def main(argv=None):
724726
legacy_text_conditioning=args.legacy_text_conditioning,
725727
hparams_from_wandb=args.hparams_file_from_wandb,
726728
phoneme_tokenizer_path=getattr(args, 'phoneme_tokenizer_path', None),
727-
disable_cas_for_context_text=args.disable_cas_for_context_text,
728729
)
729730

730731
# Load model
@@ -767,7 +768,6 @@ def main(argv=None):
767768
legacy_codebooks=args.legacy_codebooks,
768769
legacy_text_conditioning=args.legacy_text_conditioning,
769770
phoneme_tokenizer_path=getattr(args, 'phoneme_tokenizer_path', None),
770-
disable_cas_for_context_text=args.disable_cas_for_context_text,
771771
)
772772

773773
# Load model

0 commit comments

Comments
 (0)