Improve trainer UI and add GPT dependency
Browse files- app.py +86 -17
- requirements.txt +1 -0
app.py
CHANGED
|
@@ -14,6 +14,7 @@ import os
|
|
| 14 |
import shutil
|
| 15 |
import subprocess
|
| 16 |
import sys
|
|
|
|
| 17 |
from pathlib import Path
|
| 18 |
|
| 19 |
import gradio as gr
|
|
@@ -142,14 +143,65 @@ def latest_file(directory: Path, suffix: str):
|
|
| 142 |
return str(files[-1]) if files else None
|
| 143 |
|
| 144 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 145 |
def artifacts_summary():
|
| 146 |
-
|
| 147 |
-
|
|
|
|
|
|
|
| 148 |
lines = [
|
| 149 |
-
|
| 150 |
-
|
|
|
|
| 151 |
]
|
| 152 |
-
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 153 |
|
| 154 |
|
| 155 |
def dataset_prepared():
|
|
@@ -509,6 +561,10 @@ def create_ui():
|
|
| 509 |
"# 🎤 GPT-SoVITS 训练器 — 达妮娅语音\n"
|
| 510 |
"这个 Space 按当前 GPT-SoVITS 训练链路执行。先拿到 SoVITS 权重,再按需继续训练 GPT。"
|
| 511 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 512 |
|
| 513 |
with gr.Row():
|
| 514 |
with gr.Column(scale=1):
|
|
@@ -531,37 +587,50 @@ def create_ui():
|
|
| 531 |
sovits_save_every = gr.Slider(1, 5, value=1, step=1, label="每隔多少轮导出")
|
| 532 |
sovits_lr = gr.Slider(1e-5, 5e-4, value=1e-4, step=1e-5, label="学习率")
|
| 533 |
sovits_btn = gr.Button("开始 SoVITS 训练", variant="primary", size="lg")
|
| 534 |
-
sovits_log = gr.Textbox(label="SoVITS 训练日志", lines=18, interactive=False, autoscroll=True)
|
| 535 |
-
sovits_file = gr.File(label="最新 SoVITS 权重", interactive=False)
|
| 536 |
|
| 537 |
gr.Markdown("### 5. GPT 训练(可选)")
|
| 538 |
gpt_epochs = gr.Slider(1, 10, value=1, step=1, label="训练轮数")
|
| 539 |
gpt_batch = gr.Slider(1, 4, value=1, step=1, label="批次大小")
|
| 540 |
gpt_save_every = gr.Slider(1, 5, value=1, step=1, label="每隔多少轮导出")
|
| 541 |
gpt_btn = gr.Button("开始 GPT 训练", variant="secondary")
|
| 542 |
-
gpt_log = gr.Textbox(label="GPT 训练日志", lines=14, interactive=False, autoscroll=True)
|
| 543 |
-
gpt_file = gr.File(label="最新 GPT 权重", interactive=False)
|
| 544 |
|
| 545 |
-
|
| 546 |
-
|
| 547 |
-
|
| 548 |
-
|
| 549 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 550 |
|
| 551 |
env_btn.click(check_environment, outputs=env_out)
|
| 552 |
dataset_btn.click(download_dataset, outputs=dataset_out)
|
| 553 |
prep_btn.click(prepare_data, outputs=prep_out)
|
| 554 |
-
|
|
|
|
| 555 |
start_training,
|
| 556 |
inputs=[sovits_epochs, sovits_batch, sovits_save_every, sovits_lr],
|
| 557 |
outputs=[sovits_log, sovits_file],
|
| 558 |
)
|
| 559 |
-
gpt_btn.click(
|
| 560 |
start_gpt_training,
|
| 561 |
inputs=[gpt_epochs, gpt_batch, gpt_save_every],
|
| 562 |
outputs=[gpt_log, gpt_file],
|
| 563 |
)
|
| 564 |
-
refresh_btn.click(refresh_outputs, outputs=
|
|
|
|
|
|
|
|
|
|
| 565 |
|
| 566 |
return demo
|
| 567 |
|
|
|
|
| 14 |
import shutil
|
| 15 |
import subprocess
|
| 16 |
import sys
|
| 17 |
+
from datetime import datetime
|
| 18 |
from pathlib import Path
|
| 19 |
|
| 20 |
import gradio as gr
|
|
|
|
| 143 |
return str(files[-1]) if files else None
|
| 144 |
|
| 145 |
|
| 146 |
+
def list_files(directory: Path, suffix: str):
|
| 147 |
+
return sorted(directory.glob(f"*{suffix}"), key=lambda item: item.stat().st_mtime, reverse=True)
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def format_size(size_bytes: int):
|
| 151 |
+
value = float(size_bytes)
|
| 152 |
+
units = ["B", "KB", "MB", "GB"]
|
| 153 |
+
for unit in units:
|
| 154 |
+
if value < 1024 or unit == units[-1]:
|
| 155 |
+
if unit == "B":
|
| 156 |
+
return f"{int(value)} {unit}"
|
| 157 |
+
return f"{value:.1f} {unit}"
|
| 158 |
+
value /= 1024
|
| 159 |
+
return f"{int(size_bytes)} B"
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def format_mtime(path: Path):
|
| 163 |
+
return datetime.fromtimestamp(path.stat().st_mtime).strftime("%Y-%m-%d %H:%M:%S")
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def directory_overview():
|
| 167 |
+
return "\n".join(
|
| 168 |
+
[
|
| 169 |
+
f"工作目录: {WORK_DIR}",
|
| 170 |
+
f"数据集目录: {DATASET_DIR}",
|
| 171 |
+
f"日志目录: {EXP_DIR}",
|
| 172 |
+
f"SoVITS 导出目录: {SOVITS_OUTPUT_DIR}",
|
| 173 |
+
f"GPT 导出目录: {GPT_OUTPUT_DIR}",
|
| 174 |
+
]
|
| 175 |
+
)
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
def artifact_lines(label: str, files):
|
| 179 |
+
if not files:
|
| 180 |
+
return [f"{label}: 暂无"]
|
| 181 |
+
lines = [f"{label}: 共 {len(files)} 个"]
|
| 182 |
+
for index, path in enumerate(files[:10], start=1):
|
| 183 |
+
lines.append(f"{index}. {path.name} | {format_size(path.stat().st_size)} | {format_mtime(path)}")
|
| 184 |
+
return lines
|
| 185 |
+
|
| 186 |
+
|
| 187 |
def artifacts_summary():
|
| 188 |
+
sovits_files = list_files(SOVITS_OUTPUT_DIR, ".pth")
|
| 189 |
+
gpt_files = list_files(GPT_OUTPUT_DIR, ".ckpt")
|
| 190 |
+
sovits = str(sovits_files[0]) if sovits_files else None
|
| 191 |
+
gpt = str(gpt_files[0]) if gpt_files else None
|
| 192 |
lines = [
|
| 193 |
+
"训练结果总览",
|
| 194 |
+
*artifact_lines("SoVITS", sovits_files),
|
| 195 |
+
*artifact_lines("GPT", gpt_files),
|
| 196 |
]
|
| 197 |
+
return (
|
| 198 |
+
"\n".join(lines),
|
| 199 |
+
directory_overview(),
|
| 200 |
+
sovits,
|
| 201 |
+
gpt,
|
| 202 |
+
[str(item) for item in sovits_files] or None,
|
| 203 |
+
[str(item) for item in gpt_files] or None,
|
| 204 |
+
)
|
| 205 |
|
| 206 |
|
| 207 |
def dataset_prepared():
|
|
|
|
| 561 |
"# 🎤 GPT-SoVITS 训练器 — 达妮娅语音\n"
|
| 562 |
"这个 Space 按当前 GPT-SoVITS 训练链路执行。先拿到 SoVITS 权重,再按需继续训练 GPT。"
|
| 563 |
)
|
| 564 |
+
gr.Markdown(
|
| 565 |
+
"打开页面会自动加载当前输出。训练完成后,最新模型会直接出现在下载框里,"
|
| 566 |
+
"下面的“输出与目录”也会自动刷新,避免找不到文件。"
|
| 567 |
+
)
|
| 568 |
|
| 569 |
with gr.Row():
|
| 570 |
with gr.Column(scale=1):
|
|
|
|
| 587 |
sovits_save_every = gr.Slider(1, 5, value=1, step=1, label="每隔多少轮导出")
|
| 588 |
sovits_lr = gr.Slider(1e-5, 5e-4, value=1e-4, step=1e-5, label="学习率")
|
| 589 |
sovits_btn = gr.Button("开始 SoVITS 训练", variant="primary", size="lg")
|
|
|
|
|
|
|
| 590 |
|
| 591 |
gr.Markdown("### 5. GPT 训练(可选)")
|
| 592 |
gpt_epochs = gr.Slider(1, 10, value=1, step=1, label="训练轮数")
|
| 593 |
gpt_batch = gr.Slider(1, 4, value=1, step=1, label="批次大小")
|
| 594 |
gpt_save_every = gr.Slider(1, 5, value=1, step=1, label="每隔多少轮导出")
|
| 595 |
gpt_btn = gr.Button("开始 GPT 训练", variant="secondary")
|
|
|
|
|
|
|
| 596 |
|
| 597 |
+
gr.Markdown("### 6. SoVITS 实时日志与下载")
|
| 598 |
+
sovits_log = gr.Textbox(label="SoVITS 训练日志", lines=22, interactive=False, autoscroll=True)
|
| 599 |
+
sovits_file = gr.File(label="最新 SoVITS 权重", interactive=False)
|
| 600 |
+
|
| 601 |
+
gr.Markdown("### 7. GPT 实时日志与下载")
|
| 602 |
+
gpt_log = gr.Textbox(label="GPT 训练日志", lines=22, interactive=False, autoscroll=True)
|
| 603 |
+
gpt_file = gr.File(label="最新 GPT 权重", interactive=False)
|
| 604 |
+
|
| 605 |
+
gr.Markdown("### 8. 输出与目录")
|
| 606 |
+
refresh_btn = gr.Button("刷新输出与目录", variant="secondary")
|
| 607 |
+
refresh_text = gr.Textbox(label="模型列表与状态", lines=14, interactive=False, autoscroll=True)
|
| 608 |
+
output_dirs = gr.Textbox(label="工作目录与输出目录", lines=6, interactive=False)
|
| 609 |
+
with gr.Row():
|
| 610 |
+
refresh_sovits = gr.File(label="最新 SoVITS 输出", interactive=False)
|
| 611 |
+
refresh_gpt = gr.File(label="最新 GPT 输出", interactive=False)
|
| 612 |
+
with gr.Row():
|
| 613 |
+
all_sovits = gr.File(label="全部 SoVITS 文件", interactive=False, file_count="multiple")
|
| 614 |
+
all_gpt = gr.File(label="全部 GPT 文件", interactive=False, file_count="multiple")
|
| 615 |
|
| 616 |
env_btn.click(check_environment, outputs=env_out)
|
| 617 |
dataset_btn.click(download_dataset, outputs=dataset_out)
|
| 618 |
prep_btn.click(prepare_data, outputs=prep_out)
|
| 619 |
+
refresh_outputs_targets = [refresh_text, output_dirs, refresh_sovits, refresh_gpt, all_sovits, all_gpt]
|
| 620 |
+
sovits_event = sovits_btn.click(
|
| 621 |
start_training,
|
| 622 |
inputs=[sovits_epochs, sovits_batch, sovits_save_every, sovits_lr],
|
| 623 |
outputs=[sovits_log, sovits_file],
|
| 624 |
)
|
| 625 |
+
gpt_event = gpt_btn.click(
|
| 626 |
start_gpt_training,
|
| 627 |
inputs=[gpt_epochs, gpt_batch, gpt_save_every],
|
| 628 |
outputs=[gpt_log, gpt_file],
|
| 629 |
)
|
| 630 |
+
refresh_btn.click(refresh_outputs, outputs=refresh_outputs_targets)
|
| 631 |
+
sovits_event.then(refresh_outputs, outputs=refresh_outputs_targets)
|
| 632 |
+
gpt_event.then(refresh_outputs, outputs=refresh_outputs_targets)
|
| 633 |
+
demo.load(refresh_outputs, outputs=refresh_outputs_targets)
|
| 634 |
|
| 635 |
return demo
|
| 636 |
|
requirements.txt
CHANGED
|
@@ -10,6 +10,7 @@ torchaudio>=2.0.0
|
|
| 10 |
pytorch-lightning>=2.4
|
| 11 |
torchmetrics<=1.5
|
| 12 |
tensorboard
|
|
|
|
| 13 |
transformers>=4.43,<=4.50
|
| 14 |
sentencepiece>=0.1.99
|
| 15 |
accelerate>=0.20.0
|
|
|
|
| 10 |
pytorch-lightning>=2.4
|
| 11 |
torchmetrics<=1.5
|
| 12 |
tensorboard
|
| 13 |
+
matplotlib
|
| 14 |
transformers>=4.43,<=4.50
|
| 15 |
sentencepiece>=0.1.99
|
| 16 |
accelerate>=0.20.0
|