File size: 27,785 Bytes
426a190
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
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; }
"""
    )