Sfoglia il codice sorgente

refactor: 重构语音识别转录处理逻辑,拆分前后端转录处理

1. 新增TranscriptProcessor类处理流式转录的去重、拼接和分段
2. 重构前后端消息类型,拆分partial和utterance消息
3. 移除前端本地的转录拼接逻辑,改为依赖后端处理结果
4. 新增finalize消息用于强制结束当前 utterance
wenhongquan 3 settimane fa
parent
commit
173eb24fc6

+ 25 - 6
flutter_asr_client/lib/models/websocket_message.dart

@@ -31,6 +31,13 @@ final class ClearMessage extends WebSocketMessage {
   Map<String, dynamic> toJson() => {'type': 'clear'};
 }
 
+final class FinalizeMessage extends WebSocketMessage {
+  const FinalizeMessage();
+
+  @override
+  Map<String, dynamic> toJson() => {'type': 'finalize'};
+}
+
 sealed class ServerMessage {
   const ServerMessage();
 
@@ -41,10 +48,15 @@ sealed class ServerMessage {
         return ConnectedMessage(
           model: json['model'] as Map<String, dynamic>? ?? const {},
         );
-      case 'transcript':
-        return TranscriptMessage(
+      case 'partial':
+        return PartialMessage(
+          text: json['text'] as String? ?? '',
+          seq: json['seq'] as int? ?? 0,
+        );
+      case 'utterance':
+        return UtteranceMessage(
           text: json['text'] as String? ?? '',
-          language: json['language'] as String?,
+          seq: json['seq'] as int? ?? 0,
         );
       case 'cleared':
         return const ClearedMessage();
@@ -73,11 +85,18 @@ final class ConnectedMessage extends ServerMessage {
   final Map<String, dynamic> model;
 }
 
-final class TranscriptMessage extends ServerMessage {
-  const TranscriptMessage({required this.text, this.language});
+final class PartialMessage extends ServerMessage {
+  const PartialMessage({required this.text, this.seq = 0});
+
+  final String text;
+  final int seq;
+}
+
+final class UtteranceMessage extends ServerMessage {
+  const UtteranceMessage({required this.text, this.seq = 0});
 
   final String text;
-  final String? language;
+  final int seq;
 }
 
 final class ClearedMessage extends ServerMessage {

+ 9 - 73
flutter_asr_client/lib/pages/recording/notifiers/recording_notifier.dart

@@ -102,19 +102,6 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
   Duration _accumulated = Duration.zero;
   var _isSendingAudio = false;
 
-  /// Tracks the last raw transcript from server.
-  String _lastSeenText = '';
-
-  /// Accumulated full utterance, built by merging overlapping transcripts.
-  String _accumulatedText = '';
-
-  /// Consecutive times the same text was received (no ASR change).
-  int _sameCount = 0;
-
-  Timer? _utteranceTimer;
-
-  static const _utteranceSilence = Duration(seconds: 3);
-
   WebSocketService get _webSocketService => ref.read(webSocketServiceProvider);
 
   AudioCaptureService get _audioCaptureService =>
@@ -126,7 +113,6 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
       _timer?.cancel();
       _transcribeTimer?.cancel();
       _audioSendTimer?.cancel();
-      _utteranceTimer?.cancel();
     });
 
     final initialState = const RecordingState(
@@ -260,10 +246,10 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
       switch (message) {
         case ConnectedMessage():
           _updateState((s) => s.copyWith(isConnected: true));
-        case TranscriptMessage(:final text):
-          if (text.isNotEmpty) {
-            _addTranscript(text);
-          }
+        case PartialMessage(:final text):
+          _updateState((s) => s.copyWith(liveUtterance: text));
+        case UtteranceMessage(:final text):
+          _finalizeBubble(text);
         case ErrorMessage(:final message):
           _updateState((s) => s.copyWith(errorMessage: message));
         case ClearedMessage():
@@ -272,56 +258,8 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
     }, onError: (_) {});
   }
 
-  void _addTranscript(String text) {
-    if (text.isEmpty) return;
-
-    _accumulatedText = _stitch(_accumulatedText, text);
-    _updateState((s) => s.copyWith(liveUtterance: _accumulatedText));
-
-    if (text == _lastSeenText) {
-      _sameCount++;
-      if (_sameCount >= 2 && _utteranceTimer == null) {
-        // Text hasn't changed for 2+ intervals → user paused.
-        _utteranceTimer = Timer(_utteranceSilence, _finalizeUtterance);
-      }
-      return;
-    }
-    // New or corrected text → user still speaking.
-    _sameCount = 0;
-    _lastSeenText = text;
-    _utteranceTimer?.cancel();
-    _utteranceTimer = null;
-  }
-
-  /// Stitches [previous] with [next] by finding their longest suffix–prefix
-  /// overlap. If no overlap exists, [next] is a new utterance — finalize
-  /// [previous] and return [next] alone.
-  String _stitch(String previous, String next) {
-    if (previous.isEmpty) return next;
-    if (next.startsWith(previous)) return next;
-
-    // Find the longest suffix of `previous` that matches a prefix of `next`.
-    for (var n = next.length; n > 0; n--) {
-      final needle = next.substring(0, n);
-      if (previous.endsWith(needle)) {
-        return previous + next.substring(n);
-      }
-    }
-
-    // No overlap: ASR reset. Finalize old, start new.
-    _finalizeUtterance();
-    return next;
-  }
-
-  void _finalizeUtterance() {
-    _utteranceTimer?.cancel();
-    _utteranceTimer = null;
-    final text = _accumulatedText.trim();
-    _accumulatedText = '';
-    _lastSeenText = '';
-    _sameCount = 0;
-    if (text.isEmpty) return;
-
+  void _finalizeBubble(String text) {
+    if (text.trim().isEmpty) return;
     final item = ConversationItem(
       id: _uuid.v4(),
       type: ConversationItemType.userBubble,
@@ -341,7 +279,8 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
     _transcribeTimer?.cancel();
     _audioSendTimer?.cancel();
     _audioCaptureService.stopRecording();
-    _finalizeUtterance();
+    _webSocketService.send(const FinalizeMessage());
+    // Let any final utterance message from server arrive before clearing UI.
     _updateState((s) => s.copyWith(
       status: RecordingStatus.paused,
       liveUtterance: null,
@@ -374,10 +313,7 @@ final class RecordingNotifier extends AutoDisposeAsyncNotifier<RecordingState> {
     _timer?.cancel();
     _transcribeTimer?.cancel();
     _audioSendTimer?.cancel();
-    final live = state.value?.liveUtterance?.trim();
-    if (live != null && live.isNotEmpty) {
-      _finalizeUtterance();
-    }
+    _webSocketService.send(const FinalizeMessage());
     await _audioCaptureService.stopRecording();
     _webSocketService.disconnect();
     _updateState(

+ 22 - 10
whisper-qt-client/whisper_asr/server.py

@@ -14,6 +14,7 @@ from websockets.server import WebSocketServerProtocol
 sys.stdout = os.fdopen(sys.stdout.fileno(), 'w', buffering=1)
 
 from .audio_processor import AudioBuffer
+from .transcript_processor import TranscriptProcessor
 
 try:
     from .qwen_engine import QwenASREngine
@@ -44,6 +45,7 @@ class ASRServer:
         self.model_type = model_type
         self.clients: Set[WebSocketServerProtocol] = set()
         self.client_buffers: dict = {}
+        self.client_processors: dict = {}
         self.is_running = True
         
         print(f"Initializing {model_type.upper()} ASR engine...")
@@ -71,6 +73,7 @@ class ASRServer:
             max_duration=self.buffer_duration,
             sample_rate=16000
         )
+        self.client_processors[id(websocket)] = TranscriptProcessor()
         print(f"Client connected: {websocket.remote_address}")
         
         await self.send_message(websocket, {
@@ -81,6 +84,7 @@ class ASRServer:
     async def unregister(self, websocket: WebSocketServerProtocol):
         self.clients.remove(websocket)
         self.client_buffers.pop(id(websocket), None)
+        self.client_processors.pop(id(websocket), None)
         print(f"Client disconnected: {websocket.remote_address}")
     
     async def send_message(self, websocket: WebSocketServerProtocol, message: dict):
@@ -131,20 +135,17 @@ class ASRServer:
                     
         elif msg_type == "transcribe":
             client_buffer = self.client_buffers.get(id(websocket))
-            if client_buffer:
+            processor = self.client_processors.get(id(websocket))
+            if client_buffer and processor:
                 audio = client_buffer.get_last(self.buffer_duration)
-                print(f"Transcribe request: {len(audio)} samples in buffer")
                 if len(audio) > 1600:
                     try:
-                        print("Starting transcription...")
                         result = self.engine.transcribe_array(audio)
-                        print(f"Transcription result: {result['text']}")
-                        await self.send_message(websocket, {
-                            "type": "transcript",
-                            "text": result["text"],
-                            "language": result["language"],
-                            "segments": result.get("segments", [])
-                        })
+                        raw_text = result["text"].strip()
+                        if raw_text:
+                            out = processor.feed(raw_text)
+                            if out is not None:
+                                await self.send_message(websocket, out)
                     except Exception as e:
                         print(f"Transcription error: {e}")
                         import traceback
@@ -158,7 +159,18 @@ class ASRServer:
             client_buffer = self.client_buffers.get(id(websocket))
             if client_buffer:
                 client_buffer.clear()
+            processor = self.client_processors.get(id(websocket))
+            if processor:
+                processor = TranscriptProcessor()
+                self.client_processors[id(websocket)] = processor
             await self.send_message(websocket, {"type": "cleared"})
+        
+        elif msg_type == "finalize":
+            processor = self.client_processors.get(id(websocket))
+            if processor:
+                out = processor.finalize()
+                if out is not None:
+                    await self.send_message(websocket, out)
     
     async def handler(self, websocket: WebSocketServerProtocol):
         await self.register(websocket)

+ 150 - 0
whisper-qt-client/whisper_asr/transcript_processor.py

@@ -0,0 +1,150 @@
+"""
+Transcript Processor — deduplicates, stitches overlapping ASR windows,
+detects utterance boundaries, and formats streaming output for the client.
+
+Flow:
+  raw ASR transcript (sliding window, every ~2 s)
+      │
+      ▼
+  _stitch() — overlap-merge with accumulated buffer
+      │
+      ▼
+  _detect_boundary() — same text repeated = silence = utterance end
+      │
+      ├── partial → {"type": "partial", "text": accumulated, "seq": N}
+      │
+      └── (after 3× same text + 3 s silence)
+          utterance → {"type": "utterance", "text": full, "seq": N}
+"""
+
+import time
+from dataclasses import dataclass, field
+from typing import Optional, Callable, Awaitable
+
+
+@dataclass
+class _Utterance:
+    """Tracks a single utterance being built."""
+    accumulated: str = ""
+    last_raw: str = ""
+    same_count: int = 0
+    silence_since: float = 0.0
+    finalized: bool = False
+    seq: int = 0  # monotonic counter for this utterance
+
+
+class TranscriptProcessor:
+    """
+    Processes streaming ASR transcripts into clean, segmented utterances.
+
+    Pipeline (extensible):
+      1. stitch   — overlap-based deduplication
+      2. boundary — silence-based utterance segmentation
+      3. post     — (future) LLM-based punctuation/correction
+    """
+
+    def __init__(
+        self,
+        silence_repeats: int = 3,    # same-raw-text count before silence kicks in
+        silence_seconds: float = 3.0,  # timer after silence detection
+        post_hook: Optional[Callable[[str], Awaitable[str]]] = None,
+    ):
+        self.silence_repeats = silence_repeats
+        self.silence_seconds = silence_seconds
+        self._post_hook = post_hook  # LLM correction hook (future)
+
+        self._current: Optional[_Utterance] = None
+        self._global_seq = 0
+
+    # ------------------------------------------------------------------
+    # Public API
+    # ------------------------------------------------------------------
+
+    def feed(self, raw_text: str) -> dict | None:
+        """
+        Feed a raw ASR transcript. Returns a message dict to send to the
+        client, or None if nothing should be sent this tick.
+
+        Message shapes:
+          {"type": "partial",   "text": "...", "seq": N}
+          {"type": "utterance", "text": "...", "seq": N}
+        """
+        raw_text = raw_text.strip()
+        self._ensure_utterance()
+
+        # ---- 1. stitch ----
+        if self._current.accumulated:
+            merged = self._stitch(self._current.accumulated, raw_text)
+        else:
+            merged = raw_text
+
+        self._current.accumulated = merged
+
+        # ---- 2. boundary detection ----
+        if raw_text == self._current.last_raw:
+            self._current.same_count += 1
+            if self._current.same_count >= self.silence_repeats:
+                if self._current.silence_since == 0.0:
+                    self._current.silence_since = time.monotonic()
+        else:
+            self._current.same_count = 0
+            self._current.silence_since = 0.0
+            self._global_seq += 1
+            self._current.last_raw = raw_text
+            return {"type": "partial", "text": merged, "seq": self._global_seq}
+
+        # ---- 3. silence timer ----
+        if (
+            self._current.silence_since > 0
+            and time.monotonic() - self._current.silence_since >= self.silence_seconds
+        ):
+            return self._finalize()
+
+        return None
+
+    def finalize(self) -> dict | None:
+        """Force-finalize the current utterance (e.g. on pause/end)."""
+        self._ensure_utterance()
+        return self._finalize()
+
+    # ------------------------------------------------------------------
+    # Internals
+    # ------------------------------------------------------------------
+
+    def _ensure_utterance(self):
+        if self._current is None:
+            self._current = _Utterance(seq=self._global_seq)
+
+    def _stitch(self, previous: str, next_text: str) -> str:
+        """
+        Stitch two overlapping transcripts by finding the longest suffix of
+        *previous* that matches a prefix of *next_text*, then appending only
+        the non-overlapping suffix. If there is no overlap, *next_text* is a
+        fresh utterance — finalize current and return *next_text* alone.
+        """
+        if not next_text:
+            return previous
+        if next_text.startswith(previous):
+            return next_text
+
+        # Longest suffix–prefix overlap
+        for n in range(len(next_text), 0, -1):
+            needle = next_text[:n]
+            if previous.endswith(needle):
+                return previous + next_text[n:]
+
+        # No overlap → new utterance
+        self._finalize()
+        self._current = _Utterance(seq=self._global_seq)
+        return next_text
+
+    def _finalize(self) -> dict | None:
+        """Finalize current utterance and return the utterance message."""
+        if self._current is None:
+            return None
+        text = self._current.accumulated.strip()
+        self._current = None
+        self._global_seq += 1
+        if not text:
+            return None
+        return {"type": "utterance", "text": text, "seq": self._global_seq}