返回 F5-TTS
eval_infer_batch.py
根目录 / src / f5_tts / eval / eval_infer_batch.py
1 import os
2 import sys
3
4
5 sys.path.append(os.getcwd())
6
7 import argparse
8 import time
9 from importlib.resources import files
10
11 import torch
12 import torchaudio
13 from accelerate import Accelerator
14 from hydra.utils import get_class
15 from omegaconf import OmegaConf
16 from tqdm import tqdm
17
18 from f5_tts.eval.utils_eval import (
19 get_inference_prompt,
20 get_librispeech_test_clean_metainfo,
21 get_seedtts_testset_metainfo,
22 )
23 from f5_tts.infer.utils_infer import load_checkpoint, load_vocoder
24 from f5_tts.model import CFM
25 from f5_tts.model.utils import get_tokenizer
26
27
28 accelerator = Accelerator()
29 device = f"cuda:{accelerator.process_index}"
30
31
32 use_ema = True
33 target_rms = 0.1
34
35
36 rel_path = str(files("f5_tts").joinpath("../../"))
37
38
39 def main():
40 parser = argparse.ArgumentParser(description="batch inference")
41
42 parser.add_argument("-s", "--seed", default=None, type=int)
43 parser.add_argument("-n", "--expname", required=True)
44 parser.add_argument("-c", "--ckptstep", default=1250000, type=int)
45
46 parser.add_argument("-nfe", "--nfestep", default=32, type=int)
47 parser.add_argument("-o", "--odemethod", default="euler")
48 parser.add_argument("-ss", "--swaysampling", default=-1, type=float)
49
50 parser.add_argument("-t", "--testset", required=True)
51 parser.add_argument(
52 "-p", "--librispeech_test_clean_path", default=f"{rel_path}/data/LibriSpeech/test-clean", type=str
53 )
54
55 parser.add_argument("--local", action="store_true", help="Use local vocoder checkpoint directory")
56
57 args = parser.parse_args()
58
59 seed = args.seed
60 exp_name = args.expname
61 ckpt_step = args.ckptstep
62
63 nfe_step = args.nfestep
64 ode_method = args.odemethod
65 sway_sampling_coef = args.swaysampling
66
67 testset = args.testset
68
69 infer_batch_size = 1 # max frames. 1 for ddp single inference (recommended)
70 cfg_strength = 2.0
71 speed = 1.0
72 use_truth_duration = False
73 no_ref_audio = False
74
75 model_cfg = OmegaConf.load(str(files("f5_tts").joinpath(f"configs/{exp_name}.yaml")))
76 model_cls = get_class(f"f5_tts.model.{model_cfg.model.backbone}")
77 model_arc = model_cfg.model.arch
78
79 dataset_name = model_cfg.datasets.name
80 tokenizer = model_cfg.model.tokenizer
81
82 mel_spec_type = model_cfg.model.mel_spec.mel_spec_type
83 target_sample_rate = model_cfg.model.mel_spec.target_sample_rate
84 n_mel_channels = model_cfg.model.mel_spec.n_mel_channels
85 hop_length = model_cfg.model.mel_spec.hop_length
86 win_length = model_cfg.model.mel_spec.win_length
87 n_fft = model_cfg.model.mel_spec.n_fft
88
89 if testset == "ls_pc_test_clean":
90 metalst = rel_path + "/data/librispeech_pc_test_clean_cross_sentence.lst"
91 librispeech_test_clean_path = args.librispeech_test_clean_path
92 metainfo = get_librispeech_test_clean_metainfo(metalst, librispeech_test_clean_path)
93
94 elif testset == "seedtts_test_zh":
95 metalst = rel_path + "/data/seedtts_testset/zh/meta.lst"
96 metainfo = get_seedtts_testset_metainfo(metalst)
97
98 elif testset == "seedtts_test_en":
99 metalst = rel_path + "/data/seedtts_testset/en/meta.lst"
100 metainfo = get_seedtts_testset_metainfo(metalst)
101
102 # path to save genereted wavs
103 output_dir = (
104 f"{rel_path}/"
105 f"results/{exp_name}_{ckpt_step}/{testset}/"
106 f"seed{seed}_{ode_method}_nfe{nfe_step}_{mel_spec_type}"
107 f"{f'_ss{sway_sampling_coef}' if sway_sampling_coef else ''}"
108 f"_cfg{cfg_strength}_speed{speed}"
109 f"{'_gt-dur' if use_truth_duration else ''}"
110 f"{'_no-ref-audio' if no_ref_audio else ''}"
111 )
112
113 # -------------------------------------------------#
114
115 prompts_all = get_inference_prompt(
116 metainfo,
117 speed=speed,
118 tokenizer=tokenizer,
119 target_sample_rate=target_sample_rate,
120 n_mel_channels=n_mel_channels,
121 hop_length=hop_length,
122 mel_spec_type=mel_spec_type,
123 target_rms=target_rms,
124 use_truth_duration=use_truth_duration,
125 infer_batch_size=infer_batch_size,
126 )
127
128 # Vocoder model
129 local = args.local
130 if mel_spec_type == "vocos":
131 vocoder_local_path = "../checkpoints/charactr/vocos-mel-24khz"
132 elif mel_spec_type == "bigvgan":
133 vocoder_local_path = "../checkpoints/bigvgan_v2_24khz_100band_256x"
134 vocoder = load_vocoder(vocoder_name=mel_spec_type, is_local=local, local_path=vocoder_local_path)
135
136 # Tokenizer
137 vocab_char_map, vocab_size = get_tokenizer(dataset_name, tokenizer)
138
139 # Model
140 model = CFM(
141 transformer=model_cls(**model_arc, text_num_embeds=vocab_size, mel_dim=n_mel_channels),
142 mel_spec_kwargs=dict(
143 n_fft=n_fft,
144 hop_length=hop_length,
145 win_length=win_length,
146 n_mel_channels=n_mel_channels,
147 target_sample_rate=target_sample_rate,
148 mel_spec_type=mel_spec_type,
149 ),
150 odeint_kwargs=dict(
151 method=ode_method,
152 ),
153 vocab_char_map=vocab_char_map,
154 ).to(device)
155
156 ckpt_prefix = rel_path + f"/ckpts/{exp_name}/model_{ckpt_step}"
157 if os.path.exists(ckpt_prefix + ".pt"):
158 ckpt_path = ckpt_prefix + ".pt"
159 elif os.path.exists(ckpt_prefix + ".safetensors"):
160 ckpt_path = ckpt_prefix + ".safetensors"
161 else:
162 print("Loading from self-organized training checkpoints rather than released pretrained.")
163 ckpt_prefix = rel_path + f"/{model_cfg.ckpts.save_dir}/model_{ckpt_step}"
164 if os.path.exists(ckpt_prefix + ".pt"):
165 ckpt_path = ckpt_prefix + ".pt"
166 elif os.path.exists(ckpt_prefix + ".safetensors"):
167 ckpt_path = ckpt_prefix + ".safetensors"
168 else:
169 raise ValueError("The checkpoint does not exist or cannot be found in given location.")
170
171 dtype = torch.float32 if mel_spec_type == "bigvgan" else None
172 model = load_checkpoint(model, ckpt_path, device, dtype=dtype, use_ema=use_ema)
173
174 if not os.path.exists(output_dir) and accelerator.is_main_process:
175 os.makedirs(output_dir)
176
177 # start batch inference
178 accelerator.wait_for_everyone()
179 start = time.time()
180
181 with accelerator.split_between_processes(prompts_all) as prompts:
182 for prompt in tqdm(prompts, disable=not accelerator.is_local_main_process):
183 utts, ref_rms_list, ref_mels, ref_mel_lens, total_mel_lens, final_text_list = prompt
184 ref_mels = ref_mels.to(device)
185 ref_mel_lens = torch.tensor(ref_mel_lens, dtype=torch.long).to(device)
186 total_mel_lens = torch.tensor(total_mel_lens, dtype=torch.long).to(device)
187
188 # Inference
189 with torch.inference_mode():
190 generated, _ = model.sample(
191 cond=ref_mels,
192 text=final_text_list,
193 duration=total_mel_lens,
194 lens=ref_mel_lens,
195 steps=nfe_step,
196 cfg_strength=cfg_strength,
197 sway_sampling_coef=sway_sampling_coef,
198 no_ref_audio=no_ref_audio,
199 seed=seed,
200 )
201 # Final result
202 for i, gen in enumerate(generated):
203 gen = gen[ref_mel_lens[i] : total_mel_lens[i], :].unsqueeze(0)
204 gen_mel_spec = gen.permute(0, 2, 1).to(torch.float32)
205 if mel_spec_type == "vocos":
206 generated_wave = vocoder.decode(gen_mel_spec).cpu()
207 elif mel_spec_type == "bigvgan":
208 generated_wave = vocoder(gen_mel_spec).squeeze(0).cpu()
209
210 if ref_rms_list[i] < target_rms:
211 generated_wave = generated_wave * ref_rms_list[i] / target_rms
212 torchaudio.save(f"{output_dir}/{utts[i]}.wav", generated_wave, target_sample_rate)
213
214 accelerator.wait_for_everyone()
215 if accelerator.is_main_process:
216 timediff = time.time() - start
217 print(f"Done batch inference in {timediff / 60:.2f} minutes.")
218
219
220 if __name__ == "__main__":
221 main()
222
222 lines PYTHON