Files changed (1) hide show
  1. app.py +0 -287
app.py DELETED
@@ -1,287 +0,0 @@
1
- import time
2
- import uuid
3
- import os
4
- from dataclasses import dataclass, field
5
- from typing import Optional
6
-
7
- import numpy as np
8
- import gradio as gr
9
- from scipy.signal import resample_poly
10
- from faster_whisper import WhisperModel
11
-
12
- try:
13
- import webrtcvad
14
- except ImportError as e:
15
- raise ImportError("请先安装:pip install webrtcvad-wheels") from e
16
-
17
- # 可选:设置HF_TOKEN以提升下载速度(如果有token的话)
18
- # os.environ["HF_TOKEN"] = "your_huggingface_token_here"
19
-
20
- SAMPLE_RATE = 16000
21
- FRAME_MS = 30
22
- FRAME_SAMPLES = SAMPLE_RATE * FRAME_MS // 1000
23
- SILENCE_END_MS = 900
24
- SILENCE_END_FRAMES = SILENCE_END_MS // FRAME_MS
25
- MIN_CLAUSE_SEC = 0.5
26
- PARTIAL_UPDATE_SEC = 0.8 # 更快的实时更新
27
- VAD_AGGRESSIVENESS = 3
28
- ASR_LANGUAGE = "zh"
29
-
30
- print("正在加载 whisper tiny 模型...")
31
- asr_model = WhisperModel("tiny", device="cpu", compute_type="int8")
32
- print("模型加载完成")
33
-
34
-
35
- def to_16k_mono_float32(sr: int, audio: np.ndarray) -> np.ndarray:
36
- """将音频转换为16kHz单声道float32格式"""
37
- if audio.ndim > 1:
38
- audio = audio.mean(axis=1)
39
- audio = audio.astype(np.float32)
40
- if np.abs(audio).max() > 1.0:
41
- audio = audio / 32768.0
42
- if sr != SAMPLE_RATE:
43
- audio = resample_poly(audio, SAMPLE_RATE, sr).astype(np.float32)
44
- return audio
45
-
46
-
47
- def float32_to_pcm16_bytes(audio_f32: np.ndarray) -> bytes:
48
- """将float32音频转换为PCM16字节流"""
49
- clipped = np.clip(audio_f32, -1.0, 1.0)
50
- return (clipped * 32767).astype(np.int16).tobytes()
51
-
52
-
53
- class VadSegmenter:
54
- """VAD语音活动检测器"""
55
- def __init__(self, aggressiveness=VAD_AGGRESSIVENESS):
56
- self.vad = webrtcvad.Vad(aggressiveness)
57
- self.silence_frames = 0
58
- self.in_speech = False
59
- self.last_is_speech = False
60
-
61
- def feed(self, frame_bytes: bytes) -> str:
62
- """处理一帧音频,返回状态:speaking/pause_end/silence"""
63
- is_speech = self.vad.is_speech(frame_bytes, SAMPLE_RATE)
64
- self.last_is_speech = is_speech
65
- if is_speech:
66
- self.silence_frames = 0
67
- self.in_speech = True
68
- return "speaking"
69
- if self.in_speech:
70
- self.silence_frames += 1
71
- if self.silence_frames >= SILENCE_END_FRAMES:
72
- self.in_speech = False
73
- self.silence_frames = 0
74
- return "pause_end"
75
- return "speaking"
76
- return "silence"
77
-
78
-
79
- @dataclass
80
- class DebugSession:
81
- """会话状态管理"""
82
- session_id: str
83
- vad: VadSegmenter = field(default_factory=VadSegmenter)
84
- buffer: np.ndarray = field(default_factory=lambda: np.zeros(0, dtype=np.float32))
85
- leftover: np.ndarray = field(default_factory=lambda: np.zeros(0, dtype=np.float32))
86
- confirmed_text: str = "" # 已确认的文本(最终结果)
87
- partial_text: str = "" # 实时识别的文本
88
- last_hyp: str = ""
89
- last_partial_ts: float = 0.0
90
- log_lines: list = field(default_factory=list)
91
- chunk_count: int = 0
92
- speech_frame_count: int = 0
93
- silence_frame_count: int = 0
94
-
95
-
96
- def _common_prefix(a, b):
97
- """计算两个字符串的公共前缀"""
98
- n = min(len(a), len(b))
99
- i = 0
100
- while i < n and a[i] == b[i]:
101
- i += 1
102
- return a[:i]
103
-
104
-
105
- def ingest(session: Optional[DebugSession], new_chunk, caption_box):
106
- """处理音频流的主函数"""
107
- if session is None:
108
- session = DebugSession(session_id=uuid.uuid4().hex[:8])
109
- print(f"[{session.session_id}] 新会话开始")
110
-
111
- if new_chunk is None:
112
- return session, caption_box
113
-
114
- session.chunk_count += 1
115
- sr, chunk = new_chunk
116
-
117
- # 转换音频格式
118
- audio_f32 = to_16k_mono_float32(sr, chunk)
119
- session.leftover = np.concatenate([session.leftover, audio_f32])
120
-
121
- # VAD处理
122
- while len(session.leftover) >= FRAME_SAMPLES:
123
- frame = session.leftover[:FRAME_SAMPLES]
124
- session.leftover = session.leftover[FRAME_SAMPLES:]
125
- frame_bytes = float32_to_pcm16_bytes(frame)
126
- status = session.vad.feed(frame_bytes)
127
-
128
- if session.vad.last_is_speech:
129
- session.speech_frame_count += 1
130
- else:
131
- session.silence_frame_count += 1
132
-
133
- if session.vad.last_is_speech or session.vad.in_speech:
134
- session.buffer = np.concatenate([session.buffer, frame])
135
-
136
- # 检测到停顿,进行最终识别
137
- if status == "pause_end":
138
- if len(session.buffer) < SAMPLE_RATE * MIN_CLAUSE_SEC:
139
- print(f"[{session.session_id}] 触发停顿但缓冲区太短,当误判处理")
140
- session.vad.in_speech = True
141
- session.vad.silence_frames = 0
142
- else:
143
- dur = len(session.buffer) / SAMPLE_RATE
144
- print(f"[{session.session_id}] 触发停顿,缓冲区时长={dur:.2f}s,开始最终识别...")
145
- t0 = time.time()
146
- segments, info = asr_model.transcribe(
147
- session.buffer, language=ASR_LANGUAGE, task="transcribe",
148
- beam_size=5, vad_filter=True, condition_on_previous_text=False,
149
- )
150
- text = "".join(s.text for s in segments).strip()
151
- print(f"[{session.session_id}] 最终识别耗时={time.time()-t0:.2f}s, 结果='{text}'")
152
-
153
- # 更新最终文本
154
- if text:
155
- session.confirmed_text += text + " "
156
- session.partial_text = "" # 清空实时文本
157
- session.buffer = np.zeros(0, dtype=np.float32)
158
- session.last_hyp = ""
159
-
160
- # 实时识别(每0.8秒更新一次)
161
- now = time.time()
162
- if session.vad.in_speech and now - session.last_partial_ts >= PARTIAL_UPDATE_SEC:
163
- session.last_partial_ts = now
164
- if len(session.buffer) >= SAMPLE_RATE * 0.3:
165
- t0 = time.time()
166
- segments, _ = asr_model.transcribe(
167
- session.buffer, language=ASR_LANGUAGE, task="transcribe",
168
- beam_size=1, vad_filter=False, condition_on_previous_text=False,
169
- )
170
- hyp = "".join(s.text for s in segments).strip()
171
- print(f"[{session.session_id}] 实时识别: '{hyp}'")
172
- session.partial_text = hyp # 更新实时文本
173
-
174
- # 构建字幕显示(最终文本 + 实时文本)
175
- subtitle_text = session.confirmed_text
176
- if session.partial_text:
177
- if subtitle_text:
178
- subtitle_text += f"\n[实时] {session.partial_text}"
179
- else:
180
- subtitle_text = f"[实时] {session.partial_text}"
181
-
182
- # 如果没有任何文本,显示提示
183
- if not subtitle_text:
184
- subtitle_text = "等待说话..."
185
-
186
- return session, subtitle_text
187
-
188
-
189
- # 创建Gradio界面
190
- with gr.Blocks(title="实时中文字幕", css="""
191
- .subtitle-box {
192
- font-size: 24px !important;
193
- line-height: 1.6 !important;
194
- color: #2c3e50;
195
- background: #f8f9fa;
196
- padding: 20px;
197
- border-radius: 10px;
198
- min-height: 200px;
199
- border: 2px solid #3498db;
200
- }
201
- .subtitle-box:focus {
202
- border-color: #e74c3c;
203
- }
204
- .header-text {
205
- color: #2c3e50;
206
- margin-bottom: 20px;
207
- }
208
- .status-badge {
209
- display: inline-block;
210
- padding: 5px 15px;
211
- border-radius: 20px;
212
- font-weight: bold;
213
- }
214
- """) as demo:
215
- gr.Markdown(
216
- """
217
- ## 🎙️ 实时中文字幕(类似现享字幕效果)
218
- ### 使用说明:
219
- 1. 点击下方 **🎤 点击麦克风开始说话** 按钮
220
- 2. **允许浏览器访问麦克风**
221
- 3. **正常说话**,程序会实时显示识别文本(带 [实时] 标记)
222
- 4. **停顿1秒**,程序自动确认该句并添加到最终结果
223
- 5. 支持**连续多句**识别,字幕会持续累积
224
-
225
- 📌 **提示**:识别结果会实时更新,带 `[实时]` 前缀的是未确认的临时识别
226
- """
227
- )
228
-
229
- with gr.Row():
230
- with gr.Column(scale=1):
231
- audio_in = gr.Audio(
232
- sources=["microphone"],
233
- type="numpy",
234
- streaming=True,
235
- label="🎤 点击麦克风开始说话",
236
- interactive=True
237
- )
238
-
239
- # 添加控制按钮
240
- with gr.Row():
241
- clear_btn = gr.Button("🗑️ 清空字幕", variant="secondary", size="sm")
242
-
243
- with gr.Row():
244
- subtitle = gr.Textbox(
245
- label="📝 字幕显示",
246
- lines=10,
247
- elem_classes="subtitle-box",
248
- interactive=False,
249
- value="🎤 点击麦克风,开始说话..."
250
- )
251
-
252
- # 状态管理
253
- state = gr.State(None)
254
-
255
- # 清空功能
256
- def clear_subtitle(session, current_text):
257
- if session:
258
- session.confirmed_text = ""
259
- session.partial_text = ""
260
- session.buffer = np.zeros(0, dtype=np.float32)
261
- session.last_hyp = ""
262
- return session, "字幕已清空,请开始说话..."
263
-
264
- clear_btn.click(
265
- fn=clear_subtitle,
266
- inputs=[state, subtitle],
267
- outputs=[state, subtitle]
268
- )
269
-
270
- # 音频流处理
271
- audio_in.stream(
272
- fn=ingest,
273
- inputs=[state, audio_in, subtitle],
274
- outputs=[state, subtitle],
275
- stream_every=0.3, # 每0.3秒更新一次界面
276
- concurrency_limit=1 # 限制并发
277
- )
278
-
279
-
280
- if __name__ == "__main__":
281
- demo.launch(
282
- server_name="0.0.0.0",
283
- server_port=7860,
284
- share=True, # 生成公网链接
285
- debug=False,
286
- quiet=False
287
- )