File size: 13,042 Bytes
d0a9d39
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
SpeechBrain Hugging Face Spaces Demo
=====================================
基于 Gradio 的交互式演示,用于 Hugging Face Spaces 部署。

功能:
  1. 语音识别 (ASR) — 将语音转录为文字
  2. 说话人验证 — 判断两段音频是否为同一人
  3. 语音活动检测 (VAD) — 检测音频中的语音片段
  4. 说话人嵌入提取 — 提取声纹特征向量
"""

import os
import sys
import warnings

warnings.filterwarnings("ignore")

# ============================================================
# Gradio / HuggingFace Spaces 兼容性修复
# ============================================================

# --- 修复 1: gradio_client 的 boolean schema 问题 ---
# gradio_client 在处理 additionalProperties: true 时,
# 将 boolean True 作为 schema 传入,导致 "const" in schema 崩溃
try:
    import gradio_client.utils as _gc_utils

    _original_get_type = _gc_utils.get_type

    def _patched_get_type(schema):
        if isinstance(schema, bool):
            return "boolean"
        return _original_get_type(schema)

    _gc_utils.get_type = _patched_get_type

    _original_json_schema = _gc_utils._json_schema_to_python_type

    def _patched_json_schema(schema, defs):
        if isinstance(schema, bool):
            return "boolean"
        return _original_json_schema(schema, defs)

    _gc_utils._json_schema_to_python_type = _patched_json_schema
except ImportError:
    pass

# --- 修复 2: Jinja2 缓存 key 不可哈希问题 ---
# 某些 gradio 版本传给 Jinja2 的 globals 包含不可哈希 dict
try:
    import jinja2.environment as _jinja_env

    _original_load_template = _jinja_env.Environment._load_template

    def _patched_load_template(self, name, globals):
        try:
            return _original_load_template(self, name, globals)
        except TypeError:
            # 缓存 key 不可哈希时,传入空 dict(而非 None),
            # 因为 Jinja2 会将 None 替换为 self.globals,而后者也可能包含不可哈希类型
            return _original_load_template(self, name, {})

    _jinja_env.Environment._load_template = _patched_load_template
except ImportError:
    pass

# ============================================================
# Windows 兼容性修复
# ============================================================
if sys.platform == "win32":
    import types
    for _dep in ["k2", "flair", "numba", "spacy"]:
        if _dep not in sys.modules:
            m = types.ModuleType(_dep)
            m.__path__ = []
            sys.modules[_dep] = m

import torch
import gradio as gr

# ============================================================
# 模型缓存目录(HuggingFace Spaces 会自动缓存到 /data 或默认目录)
# ============================================================
MODEL_DIR = os.path.join(os.path.dirname(__file__), "pretrained_models")

_asr_model = None
_spk_model = None
_vad_model = None


def get_device():
    """获取可用设备"""
    return "cuda" if torch.cuda.is_available() else "cpu"


# ============================================================
# 模型加载(懒加载,首次使用时才下载)
# ============================================================

def get_asr_model():
    """语音识别模型 (CRDNN + RNNLM)"""
    global _asr_model
    if _asr_model is None:
        from speechbrain.inference.ASR import EncoderDecoderASR
        print("[INFO] 正在加载 ASR 模型...")
        _asr_model = EncoderDecoderASR.from_hparams(
            source="speechbrain/asr-crdnn-rnnlm-librispeech",
            savedir=os.path.join(MODEL_DIR, "asr-crdnn-rnnlm-librispeech"),
            run_opts={"device": get_device()},
        )
    return _asr_model


def get_speaker_model():
    """说话人识别模型 (ECAPA-TDNN)"""
    global _spk_model
    if _spk_model is None:
        from speechbrain.inference.speaker import SpeakerRecognition
        print("[INFO] 正在加载说话人识别模型...")
        _spk_model = SpeakerRecognition.from_hparams(
            source="speechbrain/spkrec-ecapa-voxceleb",
            savedir=os.path.join(MODEL_DIR, "spkrec-ecapa-voxceleb"),
            run_opts={"device": get_device()},
        )
    return _spk_model


def get_vad_model():
    """语音活动检测模型 (CRDNN)"""
    global _vad_model
    if _vad_model is None:
        from speechbrain.inference.VAD import VAD
        print("[INFO] 正在加载 VAD 模型...")
        _vad_model = VAD.from_hparams(
            source="speechbrain/vad-crdnn-libriparty",
            savedir=os.path.join(MODEL_DIR, "vad-crdnn-libriparty"),
            run_opts={"device": get_device()},
        )
    return _vad_model


# ============================================================
# 功能函数
# ============================================================

def transcribe_audio(audio_file):
    """语音识别:将音频转为文字"""
    if audio_file is None:
        return "⚠️ 请先上传音频文件"

    try:
        asr = get_asr_model()
        text = asr.transcribe_file(audio_file)
        return f"📝 **识别结果:**\n\n> {text}"
    except Exception as e:
        return f"❌ 识别失败:{str(e)}"


def verify_speakers(audio1, audio2):
    """说话人验证:判断两段音频是否为同一说话人"""
    if audio1 is None or audio2 is None:
        return "⚠️ 请上传两段音频文件"

    try:
        spk = get_speaker_model()
        score, prediction = spk.verify_files(audio1, audio2)

        similarity = score.item()
        is_same = "✅ 是同一说话人" if prediction else "❌ 不是同一说话人"

        if similarity > 0.5:
            level = "🟢 高"
        elif similarity > 0.0:
            level = "🟡 中"
        else:
            level = "🔴 低"

        return (
            f"## {is_same}\n\n"
            f"| 指标 | 值 |\n"
            f"|------|----|\n"
            f"| **相似度得分** | {similarity:.4f} |\n"
            f"| **置信程度** | {level} |\n"
            f"| **判定阈值** | 0.0(高于此值判定为同一人) |"
        )
    except Exception as e:
        return f"❌ 验证失败:{str(e)}"


def detect_speech(audio_file):
    """语音活动检测:标记音频中的语音段"""
    if audio_file is None:
        return "⚠️ 请先上传音频文件"

    try:
        vad = get_vad_model()
        boundaries = vad.get_speech_segments(audio_file)

        if len(boundaries) == 0:
            return "🔇 未检测到语音"

        lines = [
            "## 🎯 检测结果\n",
            f"| 片段 | 开始 | 结束 | 时长 |",
            f"|------|------|------|------|",
        ]

        total_duration = 0
        for i, (start_t, end_t) in enumerate(boundaries, 1):
            duration = end_t - start_t
            total_duration += duration
            lines.append(
                f"| {i} | {start_t:.2f}s | {end_t:.2f}s | {duration:.2f}s |"
            )

        lines.append(
            f"\n📊 **共 {len(boundaries)} 个语音片段,总时长 {total_duration:.2f}s**"
        )
        return "\n".join(lines)
    except Exception as e:
        return f"❌ 检测失败:{str(e)}"


def extract_speaker_embedding(audio_file):
    """提取说话人嵌入向量(声纹特征)"""
    if audio_file is None:
        return "⚠️ 请先上传音频文件"

    try:
        spk = get_speaker_model()
        embedding = spk.encode_file(audio_file)
        shape = tuple(embedding.shape)
        return (
            f"✅ **提取成功!**\n\n"
            f"- 嵌入向量维度:`{shape}`\n"
            f"- 前 10 个值:`{embedding[0, :10].tolist()}`\n"
            f"- 用途:说话人聚类、验证、识别等"
        )
    except Exception as e:
        return f"❌ 提取失败:{str(e)}"


# ============================================================
# Gradio 界面
# ============================================================

def create_demo():
    device = get_device()
    device_name = "GPU 🚀" if device == "cuda" else "CPU ⚡"

    with gr.Blocks(
        title="SpeechBrain Demo — 语音处理工具箱",
        theme=gr.themes.Soft(),
        css="""
        .output-markdown { font-size: 16px; }
        footer { visibility: hidden; }
        """,
    ) as demo:
        gr.Markdown(
            f"""
            # 🧠 SpeechBrain 语音处理 Demo
            ### 基于 PyTorch 的全能语音处理工具包 | 运行设备:**{device_name}**

            上传音频文件,体验语音识别、说话人验证、语音活动检测等功能。
            首次使用会自动从 HuggingFace 下载预训练模型。
            """
        )

        with gr.Tabs():
            # ── Tab 1: 语音识别 ──
            with gr.Tab("🎙️ 语音识别 (ASR)"):
                gr.Markdown("上传英文音频,自动转录为文字。支持上传文件或麦克风录音。")
                with gr.Row():
                    with gr.Column(scale=1):
                        audio_asr = gr.Audio(
                            label="上传音频",
                            type="filepath",
                            sources=["upload", "microphone"],
                        )
                        btn_asr = gr.Button("🔍 开始识别", variant="primary", size="lg")
                    with gr.Column(scale=1):
                        output_asr = gr.Markdown(value="等待上传音频...")

                btn_asr.click(
                    fn=transcribe_audio,
                    inputs=audio_asr,
                    outputs=output_asr,
                )

            # ── Tab 2: 说话人验证 ──
            with gr.Tab("👤 说话人验证"):
                gr.Markdown("上传两段音频,判断是否为同一个人说话。")
                with gr.Row():
                    with gr.Column():
                        audio_spk1 = gr.Audio(label="音频 1", type="filepath")
                    with gr.Column():
                        audio_spk2 = gr.Audio(label="音频 2", type="filepath")
                btn_spk = gr.Button("🔍 开始验证", variant="primary", size="lg")
                output_spk = gr.Markdown(value="等待上传两段音频...")

                btn_spk.click(
                    fn=verify_speakers,
                    inputs=[audio_spk1, audio_spk2],
                    outputs=output_spk,
                )

            # ── Tab 3: 语音活动检测 ──
            with gr.Tab("📊 语音检测 (VAD)"):
                gr.Markdown("检测音频中哪些时间段有语音活动。")
                with gr.Row():
                    with gr.Column(scale=1):
                        audio_vad = gr.Audio(
                            label="上传音频",
                            type="filepath",
                            sources=["upload", "microphone"],
                        )
                        btn_vad = gr.Button("🔍 开始检测", variant="primary", size="lg")
                    with gr.Column(scale=1):
                        output_vad = gr.Markdown(value="等待上传音频...")

                btn_vad.click(
                    fn=detect_speech,
                    inputs=audio_vad,
                    outputs=output_vad,
                )

            # ── Tab 4: 说话人嵌入 ──
            with gr.Tab("🧬 声纹提取"):
                gr.Markdown("提取说话人的声纹特征向量(Embedding),可用于下游任务。")
                with gr.Row():
                    with gr.Column(scale=1):
                        audio_emb = gr.Audio(label="上传音频", type="filepath")
                        btn_emb = gr.Button("🔍 提取嵌入", variant="primary", size="lg")
                    with gr.Column(scale=1):
                        output_emb = gr.Markdown(value="等待上传音频...")

                btn_emb.click(
                    fn=extract_speaker_embedding,
                    inputs=audio_emb,
                    outputs=output_emb,
                )

        gr.Markdown(
            """
            ---
            ### 📚 相关链接
            - [SpeechBrain GitHub](https://github.com/speechbrain/speechbrain)
            - [HuggingFace 模型库](https://huggingface.co/speechbrain)
            - 模型:ASR (CRDNN+RNNLM) | 说话人识别 (ECAPA-TDNN) | VAD (CRDNN)
            """
        )

    return demo


# ============================================================
# HuggingFace Spaces 入口
# ============================================================
if __name__ == "__main__":
    demo = create_demo()
    # HuggingFace Spaces 会自动注入环境变量,本地开发时绑定 0.0.0.0:7860
    server_name = os.environ.get("GRADIO_SERVER_NAME", "0.0.0.0")
    server_port = int(os.environ.get("GRADIO_SERVER_PORT", 7860))
    demo.launch(
        server_name=server_name,
        server_port=server_port,
    )