lenML/ChatTTS-Forge
301
1import gradio as gr2 3from modules.Enhancer.ResembleEnhance import unload_enhancer4from modules.models import unload_chat_tts5from modules.speaker import speaker_mgr6from modules.webui import webui_config7from modules.webui.webui_utils import get_speaker_names8 9from .ft_ui_utils import get_datasets_listfile, run_speaker_ft10from .ProcessMonitor import ProcessMonitor11 12 13class SpeakerFt:14 def __init__(self):15 self.process_monitor = ProcessMonitor()16 self.status_str = "idle"17 18 def unload_main_thread_models(self):19 unload_chat_tts()20 unload_enhancer()21 22 def run(23 self,24 batch_size: int,25 epochs: int,26 lr: str,27 train_text: bool,28 data_path: str,29 select_speaker: str = "",30 ):31 if self.process_monitor.process:32 return33 self.unload_main_thread_models()34 spk_path = None35 if select_speaker != "" and select_speaker != "none":36 select_speaker = select_speaker.split(" : ")[1].strip()37 spk = speaker_mgr.get_speaker(select_speaker)38 if spk is None:39 return ["Speaker not found"]40 spk_filename = speaker_mgr.get_speaker_filename(spk.id)41 spk_path = f"./data/speakers/{spk_filename}"42 43 command = ["python3", "-m", "modules.finetune.train_speaker"]44 command += [45 f"--batch_size={batch_size}",46 f"--epochs={epochs}",47 f"--data_path={data_path}",48 ]49 if train_text:50 command.append("--train_text")51 if spk_path:52 command.append(f"--init_speaker={spk_path}")53 54 self.status("Training process starting")55 56 self.process_monitor.start_process(command)57 58 self.status("Training started")59 60 def status(self, text: str):61 self.status_str = text62 63 def flush(self):64 stdout, stderr = self.process_monitor.get_output()65 return f"{self.status_str}\n{stdout}\n{stderr}"66 67 def clear(self):68 self.process_monitor.stdout = ""69 self.process_monitor.stderr = ""70 self.status("Logs cleared")71 72 def stop(self):73 self.process_monitor.stop_process()74 self.status("Training stopped")75 76 77def create_speaker_ft_tab(demo: gr.Blocks):78 spk_ft = SpeakerFt()79 speakers, speaker_names = get_speaker_names()80 speaker_names = ["none"] + speaker_names81 82 with gr.Row():83 with gr.Column(scale=2):84 with gr.Group():85 gr.Markdown("🎛️hparams")86 dataset_input = gr.Dropdown(87 label="Dataset", choices=get_datasets_listfile()88 )89 lr_input = gr.Textbox(label="Learning Rate", value="1e-2")90 epochs_input = gr.Slider(91 label="Epochs", value=10, minimum=1, maximum=100, step=192 )93 batch_size_input = gr.Slider(94 label="Batch Size", value=4, minimum=1, maximum=64, step=195 )96 train_text_checkbox = gr.Checkbox(label="Train text_loss", value=True)97 init_spk_dropdown = gr.Dropdown(98 label="Initial Speaker",99 choices=speaker_names,100 value="none",101 )102 103 with gr.Group():104 start_train_btn = gr.Button("Start Training")105 stop_train_btn = gr.Button("Stop Training")106 clear_train_btn = gr.Button("Clear logs")107 with gr.Column(scale=5):108 with gr.Group():109 # log110 gr.Markdown("📜logs")111 log_output = gr.Textbox(112 show_label=False, label="Log", value="", lines=20, interactive=True113 )114 115 start_train_btn.click(116 spk_ft.run,117 inputs=[118 batch_size_input,119 epochs_input,120 lr_input,121 train_text_checkbox,122 dataset_input,123 init_spk_dropdown,124 ],125 outputs=[],126 )127 stop_train_btn.click(spk_ft.stop)128 clear_train_btn.click(spk_ft.clear)129 130 if webui_config.experimental:131 demo.load(spk_ft.flush, every=1, outputs=[log_output])132 