ForcedAligner / app.py
warry's picture
Upload 3 files
426a190 verified
Raw
History Blame Contribute Delete
27.8 kB
import os
import re
import sys
import tempfile
import subprocess
import traceback
import threading
from typing import List, Tuple, Dict
# ---------- ZeroGPU:必须尽早 import spaces(在 import torch 之前)----------
# @spaces.GPU 在非 ZeroGPU 环境下是无副作用的空操作,本地/CPU 也能正常运行。
import spaces
# ---------- 持久化 HF 缓存 ----------
# 如果 Space 挂载了 Persistent Storage(固定路径 /data),把 HF 缓存指向那里,
# 这样即使 Space 重启/重建容器,已下载过的模型权重也不需要重新从 Hub 拉取。
# 如果没有挂载 Persistent Storage,则退回容器默认缓存(容器生命周期内依然只下载一次)。
_PERSIST_DIR = "/data" if os.path.isdir("/data") and os.access("/data", os.W_OK) else None
if _PERSIST_DIR:
_HF_CACHE_DIR = os.path.join(_PERSIST_DIR, "hf_cache")
os.makedirs(_HF_CACHE_DIR, exist_ok=True)
os.environ.setdefault("HF_HOME", _HF_CACHE_DIR)
os.environ.setdefault("HF_HUB_CACHE", os.path.join(_HF_CACHE_DIR, "hub"))
os.environ.setdefault("TRANSFORMERS_CACHE", os.path.join(_HF_CACHE_DIR, "hub"))
# 启用 hf_transfer 加速下载(首次下载生效,需要 requirements.txt 中安装 hf_transfer)
os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
print(f"HF 缓存目录: {os.environ.get('HF_HOME', '(未挂载 Persistent Storage,使用容器默认缓存)')}")
import torch
import soundfile as sf
import gradio as gr
# ---------- 设备检测 ----------
# 注意:在 ZeroGPU Space 中,torch.cuda.is_available() 在主进程里也会返回 True,
# 但真正的物理 GPU 只有在被 @spaces.GPU 装饰的函数被调用的瞬间才会被挂载进来。
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
print(f"运行设备: {DEVICE}")
# ---------- 导入 CTC Forced Aligner ----------
try:
import ctc_forced_aligner
from ctc_forced_aligner import (
load_audio,
load_alignment_model,
generate_emissions,
preprocess_text,
get_alignments,
)
import ctc_forced_aligner.alignment_utils as ctc_au
import ctc_forced_aligner.text_utils as ctc_tu
CTC_AVAILABLE = True
print("✅ CTC Forced Aligner 已就绪")
except ImportError:
CTC_AVAILABLE = False
print("⚠️ CTC Forced Aligner 不可用")
# ---------- Qwen3 模型封装 ----------
QWEN_AVAILABLE = False
try:
from qwen_asr import Qwen3ForcedAligner
QWEN_AVAILABLE = True
print("✅ Qwen3-ForcedAligner 已就绪")
except ImportError:
print("⚠️ Qwen3-ForcedAligner 不可用")
# ================== 模型预加载(Space 启动时只执行一次,常驻显存/内存) ==================
# ZeroGPU 的约定:模型需要在“模块根级别”创建并 .to(...)/device_map 到 cuda,
# 这样 spaces 的接管机制才能在真正拿到物理 GPU 的瞬间把它安置上去;
# 这样写同时也保证了模型只在容器启动时加载一次,而不是每次点击按钮都重新加载。
CTC_MODEL = None
CTC_TOKENIZER = None
if CTC_AVAILABLE:
try:
_ctc_dtype = torch.float16 if DEVICE == "cuda" else torch.float32
print(f"🚀 预加载 CTC 对齐模型 (设备: {DEVICE}, dtype: {_ctc_dtype}) ...")
CTC_MODEL, CTC_TOKENIZER = load_alignment_model(DEVICE, dtype=_ctc_dtype)
print("✅ CTC 对齐模型已常驻加载")
except Exception as _e:
print(f"⚠️ CTC 模型预加载失败,本次运行将禁用该模型: {_e}")
CTC_AVAILABLE = False
QWEN_MODEL = None
if QWEN_AVAILABLE:
try:
_qwen_dtype = torch.bfloat16 if DEVICE == "cuda" else torch.float32
_qwen_device_map = "cuda:0" if DEVICE == "cuda" else "cpu"
print(f"🚀 预加载 Qwen3-ForcedAligner-0.6B (设备: {DEVICE}, dtype: {_qwen_dtype}) ...")
QWEN_MODEL = Qwen3ForcedAligner.from_pretrained(
"Qwen/Qwen3-ForcedAligner-0.6B",
dtype=_qwen_dtype,
device_map=_qwen_device_map,
)
print("✅ Qwen3-ForcedAligner 已常驻加载")
except Exception as _e:
print(f"⚠️ Qwen3 模型预加载失败,本次运行将禁用该模型: {_e}")
QWEN_AVAILABLE = False
# CTC 对齐涉及对 ctc_forced_aligner 内部函数做猴子补丁(monkey patch)。
# Gradio/ZeroGPU 默认并发处理多个请求,必须加锁避免多个请求同时替换/还原全局函数、互相干扰。
_ctc_patch_lock = threading.Lock()
# ================== 核心算法 ==================
def get_pure_text_length(text: str) -> int:
"""计算纯净字符数:去除所有标点、空格、控制字符后剩余的字符数。"""
return len(re.sub(
r'[^\w一-鿿぀-ゟ゠-ヿ]',
'', str(text)
).lower())
def merge_token_timestamps_to_sentences(
token_timestamps: List[Tuple[str, float, float]],
target_sentences: List[str],
debug: bool = False
) -> List[Dict]:
"""通过字符数累计将模型输出的词/字级时间戳匹配到预分段短句。"""
if not target_sentences:
return []
results = []
token_idx = 0
total_tokens = len(token_timestamps)
for sent_idx, sentence in enumerate(target_sentences):
t_len = get_pure_text_length(sentence)
if t_len == 0:
results.append({"text": sentence, "start": 0.0, "end": 0.0})
continue
acc_len = 0
st, et = None, None
while token_idx < total_tokens and acc_len < t_len:
seg_text, seg_start, seg_end = token_timestamps[token_idx]
if st is None:
st = seg_start
et = seg_end
acc_len += get_pure_text_length(seg_text)
token_idx += 1
if debug and sent_idx < 5:
print(f" [{sent_idx}] \"{sentence[:50]}\" -> "
f"tokens char_cnt={acc_len}/{t_len} "
f"time={st:.2f}s-{et:.2f}s " if st else " ")
if st is not None and et is not None:
results.append({
"text": sentence,
"start": round(st, 3),
"end": round(et, 3),
})
else:
results.append({"text": sentence, "start": 0.0, "end": 0.0})
# 后处理:修复缺失/异常时间戳
for i in range(len(results)):
if results[i]["start"] == 0.0 and results[i]["end"] == 0.0:
for j in range(i - 1, -1, -1):
if results[j]["end"] > 0:
results[i]["start"] = results[j]["end"]
results[i]["end"] = results[j]["end"]
break
if results[i]["start"] == 0.0:
for j in range(i + 1, len(results)):
if results[j]["start"] > 0:
results[i]["start"] = results[j]["start"]
results[i]["end"] = results[j]["start"]
break
for i in range(len(results)):
if i > 0 and results[i]["start"] < results[i - 1]["end"]:
results[i]["start"] = results[i - 1]["end"]
if results[i]["end"] < results[i]["start"]:
results[i]["end"] = results[i]["start"] + 0.001
if debug:
non_zero = sum(1 for r in results if r["start"] > 0 or r["end"] > 0)
print(f"时间戳覆盖率: {non_zero}/{len(results)} 句")
return results
def seconds_to_srt_time(seconds: float) -> str:
seconds = max(0, seconds)
h = int(seconds // 3600)
m = int((seconds % 3600) // 60)
s = int(seconds % 60)
ms = int((seconds % 1) * 1000)
return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"
def format_srt(segments: List[Dict]) -> str:
lines = []
index = 1
for seg in segments:
text = seg["text"].strip()
if not text:
continue
lines.append(str(index))
lines.append(
f"{seconds_to_srt_time(seg['start'])} --> {seconds_to_srt_time(seg['end'])}"
)
lines.append(text)
lines.append("")
index += 1
return "\n".join(lines)
# ================== SRT 时间轴二次微调(移植自 srt-time.py) ==================
def adjust_srt_timeline(segments: List[Dict], offset: float = 0.2) -> List[Dict]:
"""
对对齐后的 segments 做时间轴二次微调,逻辑与 srt-time.py 完全一致:
a. 最高原则:相邻字幕不能有交叉
b. 每条字幕开始时间提前 offset(0.2 秒),且不能小于 0
若提前后与上一条交叉(或相隔不到 0.2 秒),则将本条开始时间
收回到上一条结束时间
c. 否则把上一条结束时间向后延长 offset(0.2 秒),
但不能晚于本条(提前后的)开始时间
"""
if not segments:
return segments
# 深拷贝,避免污染原数据
adjusted = [
{
"text": seg["text"],
"start": float(seg["start"]),
"end": float(seg["end"]),
}
for seg in segments
]
for i in range(len(adjusted)):
# b. 每个序号开始时间提前 0.2 秒
adjusted[i]["start"] -= offset
# 安全边界:开始时间不能小于 0
if adjusted[i]["start"] < 0:
adjusted[i]["start"] = 0.0
if i > 0:
prev_end = adjusted[i - 1]["end"]
curr_start = adjusted[i]["start"]
# a. 最高原则:不能有交叉
if curr_start < prev_end:
# b. 提前0.2秒后如果和前一个交叉了,则提前到和前一个结束时间相等
adjusted[i]["start"] = prev_end
else:
# c. 如果仍有间隔,把上一个序号结束时间向后延长 0.2 秒
# 前提是不能晚于当前序号开始时间(提前后的)。
# 若间隔不够 0.2 秒,则延长至相等。
adjusted[i - 1]["end"] = min(prev_end + offset, curr_start)
return adjusted
# ================== CTC 对齐封装(含容错补丁) ==================
def run_ctc_alignment(
audio_path: str,
full_text: str,
target_sentences: List[str],
language: str = "eng"
) -> List[Dict]:
"""使用 CTC Forced Aligner 进行强制对齐(原补丁保留,模型已在启动时常驻加载)"""
global CTC_MODEL, CTC_TOKENIZER
_original_get_spans = ctc_au.get_spans
_original_postprocess = ctc_tu.postprocess_results
def _relaxed_get_spans(tokens_starred, segments, blank_token):
n_seg = len(segments)
spans = []
si = 0
for token in tokens_starred:
target_letters = token.split(" ")
while si < n_seg and segments[si].label == blank_token:
si += 1
start_seg_idx = si
end_seg_idx = si
matched_any = False
for ltr in target_letters:
while si < n_seg and segments[si].label == blank_token:
si += 1
if si < n_seg and segments[si].label == ltr:
if not matched_any:
start_seg_idx = si
end_seg_idx = si
matched_any = True
si += 1
if not matched_any:
safe_idx = min(start_seg_idx, n_seg - 1) if n_seg > 0 else 0
spans.append([ctc_au.Segment(token, safe_idx, safe_idx)])
else:
spans.append(segments[start_seg_idx : end_seg_idx + 1])
return spans
def _safe_postprocess_results(text_starred, spans, stride, scores, merge_threshold=0.0):
results = []
for i, t in enumerate(text_starred):
if t == "<star>": continue
span = spans[i]
if not span: continue
seg_start_idx = span[0].start
seg_end_idx = span[-1].end
audio_start_sec = seg_start_idx * stride / 1000.0
audio_end_sec = seg_end_idx * stride / 1000.0
score = scores[seg_start_idx : seg_end_idx + 1].sum() if seg_end_idx >= seg_start_idx else 0.0
score_val = score.item() if hasattr(score, "item") else float(score)
results.append({
"start": audio_start_sec,
"end": audio_end_sec,
"text": t,
"score": score_val,
})
ctc_tu.merge_segments(results, merge_threshold)
return results
# 全局猴子补丁 + 全局预加载模型都是共享状态,多个并发请求必须串行访问这一段
with _ctc_patch_lock:
try:
ctc_au.get_spans = _relaxed_get_spans
ctc_tu.postprocess_results = _safe_postprocess_results
alignment_model, alignment_tokenizer = CTC_MODEL, CTC_TOKENIZER
print("🔄 加载音频...")
audio_waveform = load_audio(audio_path, alignment_model.dtype, alignment_model.device)
print("🔄 生成发射矩阵...")
emissions, stride = generate_emissions(alignment_model, audio_waveform, batch_size=8)
non_latin = {"cmn", "zho", "chi", "jpn", "ja", "kor", "ko", "ara", "ar", "rus", "ru"}
needs_romanize = language in non_latin
tokens_starred, text_starred = preprocess_text(full_text, romanize=needs_romanize, language=language)
print("🔄 CTC 解码...")
segments_raw, scores, blank_token = get_alignments(emissions, tokens_starred, alignment_tokenizer)
print("🔄 获取时间跨度 (容错模式)...")
spans = ctc_au.get_spans(tokens_starred, segments_raw, blank_token)
results = ctc_tu.postprocess_results(text_starred, spans, stride, scores)
token_timestamps = [(seg["text"], seg["start"], seg["end"]) for seg in results]
print(f"模型输出 {len(token_timestamps)} 个词/字级时间戳")
segments = merge_token_timestamps_to_sentences(token_timestamps, target_sentences, debug=True)
finally:
# 无论成功与否都要还原全局函数,避免污染下一次请求
ctc_au.get_spans = _original_get_spans
ctc_tu.postprocess_results = _original_postprocess
# 模型是启动时预加载的全局单例,这里不再 del,只清理本次推理产生的显存碎片
if DEVICE == "cuda":
torch.cuda.empty_cache()
return segments
# ================== Qwen3 对齐封装 ==================
def run_qwen_alignment(
audio_path: str,
full_text: str,
target_sentences: List[str],
language: str = "Chinese"
) -> List[Dict]:
"""
使用 Qwen3-ForcedAligner-0.6B 进行强制对齐。
模型已在 Space 启动时常驻加载(CTC_MODEL/QWEN_MODEL 全局单例),此处直接复用,
不再每次请求都重新 from_pretrained,避免重复下载/重复加载开销。
"""
model = QWEN_MODEL
# 读取音频
audio_data, sr = sf.read(audio_path)
total_duration = len(audio_data) / sr
print(f"📊 音频总时长: {total_duration:.1f}s")
# 切片参数
MAX_CHUNK_DUR = 240.0 # 每次最多 4 分钟
SAFE_TAIL_MARGIN = 15.0 # 丢弃末尾 15s 的不完整句子
remaining = list(target_sentences)
time_offset = 0.0
all_segments = []
chunk_idx = 0
while remaining and time_offset < total_duration:
chunk_idx += 1
chunk_dur = min(MAX_CHUNK_DUR, total_duration - time_offset)
is_last = (time_offset + chunk_dur >= total_duration - 1.0)
start_f = int(time_offset * sr)
end_f = int((time_offset + chunk_dur) * sr)
chunk_audio = audio_data[start_f:end_f]
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
sf.write(f.name, chunk_audio, sr)
chunk_path = f.name
chunk_text = " ".join(remaining)
print(f"\n▶️ Chunk {chunk_idx}: 音频[{time_offset:.0f}s-{time_offset + chunk_dur:.0f}s] "
f"剩余{len(remaining)}句")
results = model.align(audio=chunk_path, text=chunk_text, language=language)
tokens = results[0] # List[AlignmentResult]
token_data = []
for seg in tokens:
try:
token_data.append((seg.text, seg.start_time, seg.end_time))
except AttributeError:
d = vars(seg) if hasattr(seg, '__dict__') else {}
token_data.append((
d.get('text', d.get('token', d.get('word', ''))),
d.get('start_time', d.get('start', 0.0)),
d.get('end_time', d.get('end', 0.0)),
))
# 用字符计数法匹配句子
matched = []
ti = 0
for sentence in remaining:
t_len = get_pure_text_length(sentence)
if t_len == 0:
continue
acc = 0
st, et = None, None
while ti < len(token_data) and acc < t_len:
seg_text, seg_start, seg_end = token_data[ti]
if st is None:
st = seg_start
et = seg_end
acc += get_pure_text_length(seg_text)
ti += 1
if st is not None and et is not None:
matched.append({"text": sentence, "start": st, "end": et})
# 安全切分点
if is_last:
valid = matched
remaining = []
else:
valid_idx = -1
for i, m in enumerate(matched):
if m["end"] < (chunk_dur - SAFE_TAIL_MARGIN):
valid_idx = i
else:
break
if valid_idx == -1 and matched:
valid_idx = 0
valid = matched[:valid_idx + 1] if valid_idx >= 0 else []
remaining = remaining[valid_idx + 1:] if valid_idx >= 0 else []
print(f" 本段对齐 {len(valid)} 句(共{len(matched)}句匹配)")
for m in valid:
all_segments.append({
"text": m["text"],
"start": round(m["start"] + time_offset, 3),
"end": round(m["end"] + time_offset, 3),
})
if valid:
time_offset = time_offset + valid[-1]["end"]
else:
time_offset = total_duration
os.unlink(chunk_path)
if DEVICE == "cuda":
torch.cuda.empty_cache()
# 模型是启动时预加载的全局单例,这里不再 del,只清理本次推理产生的显存碎片
if DEVICE == "cuda":
torch.cuda.empty_cache()
print(f"\n✅ Qwen 对齐完成:{len(all_segments)} 句")
return all_segments
# ================== 音频格式转换 ==================
def convert_to_wav(input_audio_path: str) -> str:
"""使用 ffmpeg 转换为 16kHz 单声道 wav"""
tmp_wav = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
tmp_wav.close()
cmd = [
"ffmpeg", "-y",
"-i", input_audio_path,
"-ar", "16000",
"-ac", "1",
"-c:a", "pcm_s16le",
"-loglevel", "error",
tmp_wav.name
]
try:
subprocess.run(cmd, check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
return tmp_wav.name
except subprocess.CalledProcessError as e:
raise RuntimeError(f"FFmpeg 转换失败: {e.stderr.decode('utf-8', errors='ignore')}")
# ================== ZeroGPU 时长估算 ==================
def _estimate_gpu_duration(audio_file, text_input, text_file, language, model_choice) -> int:
"""
根据音频时长粗略估算这次请求需要占用 GPU 多久(秒)。
ZeroGPU 按 @spaces.GPU(duration=...) 请求 GPU 时间片,预留过短会导致任务被中断,
预留过长则会更快消耗当天的免费 GPU 配额,这里给一个比较宽松但有上限的估算。
"""
default_duration = 60
if not audio_file:
return default_duration
try:
info = sf.info(audio_file)
audio_seconds = info.frames / float(info.samplerate)
except Exception:
return default_duration
# CTC 一次性跑完整段音频,Qwen3 按 4 分钟分片跑,两者都留一些解码/开销余量
factor = 0.5 if model_choice == "CTC Forced Aligner" else 1.0
estimated = int(audio_seconds * factor) + 30
# 60-280 秒的范围内取值;若单个账号的 ZeroGPU 时长上限不同,可按需调整上限
return max(60, min(estimated, 280))
# ================== 主处理函数 ==================
@spaces.GPU(duration=_estimate_gpu_duration)
def process_alignment(
audio_file,
text_input: str,
text_file,
language: str,
model_choice: str
):
debug_lines = []
if audio_file is None:
return "", "请上传音频文件", "", None
# 读取文本
raw_text = ""
if text_file is not None:
try:
file_path = text_file if isinstance(text_file, str) else (
text_file.get("name", "") if isinstance(text_file, dict) else getattr(text_file, "name", "")
)
if file_path and os.path.exists(file_path):
with open(file_path, "r", encoding="utf-8") as f:
raw_text = f.read()
debug_lines.append(f"从文件读取文本 ({len(raw_text)} 字符)")
except Exception as e:
debug_lines.append(f"读取文本文件失败: {e}")
if not raw_text and text_input:
raw_text = text_input
if not raw_text or not raw_text.strip():
return "", "请输入文本或上传文本文件", "", None
target_sentences = [line.strip() for line in raw_text.strip().splitlines() if line.strip()]
if not target_sentences:
return "", "文本为空或格式不正确(每行一个短句)", "", None
full_text = " ".join(target_sentences)
lang_map = {
"中文": "cmn", "英文": "eng", "日语": "jpn",
"韩语": "kor", "法语": "fra", "德语": "deu",
"俄语": "rus", "西班牙语": "spa", "意大利语": "ita",
"葡萄牙语": "por",
}
lang = lang_map.get(language, "cmn")
# Qwen 模型的语言映射(将 UI 的中文选项映射为模型需要的英文标识)
qwen_lang_map = {
"中文": "Chinese",
"英文": "English",
"日语": "Japanese",
"韩语": "Korean",
"法语": "French",
"德语": "German",
"俄语": "Russian",
"西班牙语": "Spanish",
"意大利语": "Italian",
"葡萄牙语": "Portuguese",
}
# 如果选择了不支持的语言,默认回退到 English (或 Chinese,视 Qwen3 模型的具体支持情况而定)
qwen_lang = qwen_lang_map.get(language, "English")
debug_lines.append(f"音频: {audio_file}")
debug_lines.append(f"语言: {language} (内部代码: {lang})")
debug_lines.append(f"选用模型: {model_choice}")
debug_lines.append(f"句子数: {len(target_sentences)}")
# 音频转换
try:
processed_audio_path = convert_to_wav(audio_file)
debug_lines.append("✅ 音频格式转换完成")
except Exception as e:
debug_lines.append(f"❌ 音频转码失败: {e}")
return "", "音频转码失败,请上传有效文件", "\n".join(debug_lines), None
# 选择模型执行对齐
try:
if model_choice == "CTC Forced Aligner":
if not CTC_AVAILABLE:
return "", "CTC 模型未安装,请检查依赖。", "\n".join(debug_lines), None
segments = run_ctc_alignment(processed_audio_path, full_text, target_sentences, lang)
else: # Qwen3
if not QWEN_AVAILABLE:
return "", "Qwen3 模型未安装,请检查依赖。", "\n".join(debug_lines), None
segments = run_qwen_alignment(processed_audio_path, full_text, target_sentences, qwen_lang)
os.unlink(processed_audio_path)
# ============ SRT 时间轴二次微调(集成自 srt-time.py) ============
segments = adjust_srt_timeline(segments, offset=0.2)
debug_lines.append("✅ SRT 时间轴二次微调完成(提前 0.2s / 消除交叉 / 必要时延长上一段)")
srt_content = format_srt(segments)
debug_lines.append(f"\n🎉 对齐完成! 共 {len(segments)} 段")
for seg in segments[:15]:
debug_lines.append(f" [{seg['start']:.2f}s - {seg['end']:.2f}s] {seg['text'][:60]}")
if len(segments) > 15:
debug_lines.append(f" ... 共 {len(segments)} 段")
# ================= 修改部分:生成同名 SRT 文件 =================
audio_basename = os.path.basename(audio_file)
srt_filename = os.path.splitext(audio_basename)[0] + ".srt"
srt_full_path = os.path.join(tempfile.gettempdir(), srt_filename)
with open(srt_full_path, "w", encoding="utf-8") as f:
f.write(srt_content)
# ===============================================================
return srt_content, f"对齐完成! 共 {len(segments)} 段", "\n".join(debug_lines), srt_full_path
except Exception as e:
error_detail = traceback.format_exc()
debug_lines.append(f"\n❌ 错误: {e}\n{error_detail}")
if os.path.exists(processed_audio_path):
os.unlink(processed_audio_path)
return "", f"处理出错: {str(e)}", "\n".join(debug_lines), None
# ================== Gradio 界面 ==================
with gr.Blocks(title="字幕自动打轴工具(双模型)") as demo:
gr.Markdown("""
# 字幕自动打轴工具(支持双模型)
将音频与文本自动对齐,生成带精准时间轴的 SRT 字幕文件。
""")
with gr.Row():
with gr.Column(scale=2):
audio_input = gr.Audio(label="音频文件", type="filepath")
text_input = gr.Textbox(
label="文本内容(每行一个短句)",
placeholder="今天天气真好。\n我们一起去公园吧。",
lines=8, max_lines=20
)
text_file = gr.File(label="或上传文本文件 (.txt)", file_types=[".txt"])
language_choice = gr.Dropdown(
label="音频语言",
choices=["中文", "英文", "日语", "韩语", "法语", "德语", "俄语", "西班牙语", "意大利语", "葡萄牙语"],
value="英文"
)
model_choice = gr.Dropdown(
label="对齐模型",
choices=["CTC Forced Aligner", "Qwen3-ForcedAligner-0.6B"],
value="Qwen3-ForcedAligner-0.6B"
)
submit_btn = gr.Button("开始对齐", variant="primary")
status_output = gr.Textbox(label="状态", interactive=False)
with gr.Column(scale=2):
srt_output = gr.Textbox(
label="生成的 SRT 字幕",
lines=18, max_lines=30, interactive=False,
elem_classes=["srt-output"]
)
srt_download = gr.File(label="下载 SRT 文件", interactive=False)
with gr.Accordion("调试信息", open=False):
debug_output = gr.Textbox(label="详细日志", lines=12, interactive=False)
submit_btn.click(
fn=process_alignment,
inputs=[audio_input, text_input, text_file, language_choice, model_choice],
outputs=[srt_output, status_output, debug_output, srt_download]
)
if __name__ == "__main__":
demo.queue(max_size=5).launch(
server_name="0.0.0.0",
server_port=7860,
share=False,
css="""
.srt-output textarea { font-family: "Courier New", monospace; font-size: 13px; }
footer { visibility: hidden; }
"""
)