| 32 | logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') |
| 33 | |
| 34 | def get_args(): |
| 35 | parser = argparse.ArgumentParser(description='inference only with your model') |
| 36 | parser.add_argument('--config', required=True, help='config file') |
| 37 | parser.add_argument('--prompt_data', required=True, help='prompt data file') |
| 38 | parser.add_argument('--flow_model', default=None, required=False, help='flow model file') |
| 39 | parser.add_argument('--llm_model', default=None,required=False, help='flow model file') |
| 40 | parser.add_argument('--music_tokenizer', required=True, help='music tokenizer model file') |
| 41 | parser.add_argument('--wavtokenizer', required=True, help='wavtokenizer model file') |
| 42 | parser.add_argument('--chorus', default="random",required=False, help='chorus tag generation mode, eg. random, verse, chorus, intro.') |
| 43 | parser.add_argument('--fast', action='store_true', required=False, help='True: fast inference mode, without flow matching for fast inference. False: normal inference mode, with flow matching for high quality.') |
| 44 | parser.add_argument('--batch', action='store_true', required=False, help='batch inference mode, use infer.yaml to set batch size') |
| 45 | parser.add_argument('--dtype', type=str, default="fp16", required=False, choices=["fp16", "bf16", "fp32"], help='data type') |
| 46 | parser.add_argument('--fp16', default=True, type=bool, required=False, help='inference with fp16 model') |
| 47 | parser.add_argument('--fade_out', default=True, type=bool, required=False, help='add fade out effect to generated audio') |
| 48 | parser.add_argument('--fade_out_duration', default=1.0, type=float, required=False, help='fade out duration in seconds') |
| 49 | parser.add_argument('--trim', default=False, type=bool, required=False, help='trim the silence ending of generated audio') |
| 50 | parser.add_argument('--format', type=str, default="wav", required=False, |
| 51 | choices=["wav", "mp3", "m4a", "flac"], |
| 52 | help='sampling rate of input audio') |
| 53 | parser.add_argument('--sample_rate', type=int, default=24000, required=False, |
| 54 | help='sampling rate of input audio') |
| 55 | parser.add_argument('--output_sample_rate', type=int, default=48000, required=False, choices=[24000, 48000], |
| 56 | help='sampling rate of generated output audio') |
| 57 | parser.add_argument('--min_generate_audio_seconds', type=float, default=0.0, required=False, |
| 58 | help='the minimum generated audio length in seconds') |
| 59 | parser.add_argument('--max_generate_audio_seconds', type=float, default=30.0, required=False, |
| 60 | help='the maximum generated audio length in seconds') |
| 61 | parser.add_argument('--gpu', |
| 62 | type=int, |
| 63 | default=0, |
| 64 | help='gpu id for this rank, -1 for cpu') |
| 65 | parser.add_argument('--task', |
| 66 | default='text-to-music', |
| 67 | choices=['text-to-music', 'continuation', "reconstruct", "super_resolution"], |
| 68 | help='choose inference task type. text-to-music: text-to-music task. continuation: music continuation task. reconstruct: reconstruction of original music. super_resolution: convert original 24kHz music into 48kHz music.') |
| 69 | parser.add_argument('--result_dir', required=True, help='asr result file') |
| 70 | args = parser.parse_args() |
| 71 | print(args) |
| 72 | return args |
| 73 | |
| 74 | def main(): |
| 75 | args = get_args() |