CoolFace
Apppublic

jhansss/SingingSDS

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
cli.py74 linesDownload Raw Back to root
1from argparse import ArgumentParser2from logging import getLogger3from pathlib import Path4 5import yaml6 7from characters import get_character8from pipeline import SingingDialoguePipeline9 10logger = getLogger(__name__)11 12 13def get_parser():14    parser = ArgumentParser()15    parser.add_argument("--query_audios", nargs="+", type=Path, required=True)16    parser.add_argument(17        "--config_path", type=Path, default="config/cli/yaoyin_default.yaml"18    )19    parser.add_argument("--output_audio_folder", type=Path, required=True)20    parser.add_argument("--eval_results_csv", type=Path, required=True)21    return parser22 23 24def load_config(config_path: Path):25    with open(config_path, "r") as f:26        config = yaml.safe_load(f)27    return config28 29 30def main():31    parser = get_parser()32    args = parser.parse_args()33    config = load_config(args.config_path)34    pipeline = SingingDialoguePipeline(config)35    speaker = config["speaker"]36    language = config["language"]37    character_name = config["prompt_template_character"]38    character = get_character(character_name)39    prompt_template = character.prompt40    args.output_audio_folder.mkdir(parents=True, exist_ok=True)41    args.eval_results_csv.parent.mkdir(parents=True, exist_ok=True)42    with open(args.eval_results_csv, "a") as f:43        f.write(44            f"query_audio,asr_model,llm_model,svs_model,melody_source,language,speaker,output_audio,asr_text,llm_text,metrics\n"45        )46    try:47        for query_audio in args.query_audios:48            output_audio = args.output_audio_folder / f"{query_audio.stem}_response.wav"49            results = pipeline.run(50                query_audio,51                language,52                prompt_template,53                speaker,54                output_audio_path=output_audio,55            )56            metrics = pipeline.evaluate(output_audio, **results)57            metrics.update(results.get("metrics", {}))58            metrics_str = ",".join([f"{metrics[k]}" for k in sorted(metrics.keys())])59            logger.info(60                f"Input: {query_audio}, Output: {output_audio}, ASR results: {results['asr_text']}, LLM results: {results['llm_text']}"61            )62            with open(args.eval_results_csv, "a") as f:63                f.write(64                    f"{query_audio},{config['asr_model']},{config['llm_model']},{config['svs_model']},{config['melody_source']},{config['language']},{config['speaker']},{output_audio},{results['asr_text']},{results['llm_text']},{metrics_str}\n"65                )66    except Exception as e:67        logger.error(f"Error in main: {e}")68        breakpoint()69        raise e70 71 72if __name__ == "__main__":73    main()74