jhansss/SingingSDS
0
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 