huanx commited on
Commit
0cc59b4
·
verified ·
1 Parent(s): 42cb380

Improve trainer UI and add GPT dependency

Browse files
Files changed (2) hide show
  1. app.py +86 -17
  2. 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
- sovits = latest_file(SOVITS_OUTPUT_DIR, ".pth")
147
- gpt = latest_file(GPT_OUTPUT_DIR, ".ckpt")
 
 
148
  lines = [
149
- f"SoVITS: {sovits or '暂无'}",
150
- f"GPT: {gpt or '暂无'}",
 
151
  ]
152
- return "\n".join(lines), sovits, gpt
 
 
 
 
 
 
 
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
- gr.Markdown("### 6. 当前输出")
546
- refresh_btn = gr.Button("刷新最新权重")
547
- refresh_text = gr.Textbox(label="输出摘要", lines=3, interactive=False)
548
- refresh_sovits = gr.File(label="SoVITS 输出", interactive=False)
549
- refresh_gpt = gr.File(label="GPT 输出", interactive=False)
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- sovits_btn.click(
 
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=[refresh_text, refresh_sovits, refresh_gpt])
 
 
 
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