CoolFace
Apppublic

lenML/ChatTTS-Forge

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
301likes
speaker_ft_tab.py132 linesDownload Raw Back to finetune
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