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 == "": 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; } """ )