Spaces:
Running on Zero
Running on Zero
| 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)) | |
| # ================== 主处理函数 ================== | |
| 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; } | |
| """ | |
| ) |