engine.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194
  1. """
  2. Whisper ASR Engine - Wrapper for faster-whisper
  3. """
  4. import io
  5. import numpy as np
  6. from typing import Optional, List, Callable
  7. from faster_whisper import WhisperModel
  8. class WhisperEngine:
  9. """
  10. Faster-Whisper based ASR engine with streaming support.
  11. """
  12. def __init__(
  13. self,
  14. model_size: str = "base",
  15. device: str = "auto",
  16. compute_type: str = "auto",
  17. download_root: Optional[str] = None,
  18. language: Optional[str] = "zh"
  19. ):
  20. """
  21. Initialize Whisper engine.
  22. Args:
  23. model_size: Model size (tiny, base, small, medium, large)
  24. device: Device to use (cpu, cuda, auto)
  25. compute_type: Compute type (float16, int8, auto)
  26. download_root: Model download directory
  27. language: Target language code (zh, en, etc.)
  28. """
  29. self.model_size = model_size
  30. self.language = language
  31. # Auto-detect device
  32. if device == "auto":
  33. device = "cuda" if self._check_cuda() else "cpu"
  34. # Auto-detect compute type
  35. if compute_type == "auto":
  36. compute_type = "float16" if device == "cuda" else "int8"
  37. self.device = device
  38. self.compute_type = compute_type
  39. print(f"Loading Whisper model: {model_size} on {device} ({compute_type})")
  40. self.model = WhisperModel(
  41. model_size,
  42. device=device,
  43. compute_type=compute_type,
  44. download_root=download_root
  45. )
  46. print(f"Model loaded successfully")
  47. def _check_cuda(self) -> bool:
  48. """Check if CUDA is available."""
  49. try:
  50. import torch
  51. return torch.cuda.is_available()
  52. except ImportError:
  53. return False
  54. def transcribe_audio(
  55. self,
  56. audio_data: bytes,
  57. sample_rate: int = 16000,
  58. callback: Optional[Callable] = None
  59. ) -> dict:
  60. """
  61. Transcribe audio data.
  62. Args:
  63. audio_data: Raw PCM audio bytes (16-bit, mono)
  64. sample_rate: Audio sample rate
  65. callback: Optional callback for streaming results
  66. Returns:
  67. Dictionary with transcription result
  68. """
  69. # Convert bytes to numpy array
  70. audio_array = self._bytes_to_array(audio_data)
  71. # Resample if needed
  72. if sample_rate != 16000:
  73. audio_array = self._resample(audio_array, sample_rate, 16000)
  74. # Run inference
  75. segments, info = self.model.transcribe(
  76. audio_array,
  77. language=self.language,
  78. beam_size=5,
  79. vad_filter=True,
  80. vad_parameters=dict(min_silence_duration_ms=500)
  81. )
  82. # Collect results
  83. full_text = ""
  84. segment_list = []
  85. for segment in segments:
  86. text = segment.text.strip()
  87. segment_list.append({
  88. 'start': segment.start,
  89. 'end': segment.end,
  90. 'text': text
  91. })
  92. full_text += text + " "
  93. if callback:
  94. callback({
  95. 'text': text,
  96. 'start': segment.start,
  97. 'end': segment.end,
  98. 'is_final': False
  99. })
  100. return {
  101. 'text': full_text.strip(),
  102. 'segments': segment_list,
  103. 'language': info.language,
  104. 'language_probability': info.language_probability
  105. }
  106. def transcribe_array(
  107. self,
  108. audio_array: np.ndarray,
  109. sample_rate: int = 16000
  110. ) -> dict:
  111. """
  112. Transcribe from numpy array.
  113. Args:
  114. audio_array: Audio as numpy array (float32, range [-1, 1])
  115. sample_rate: Audio sample rate
  116. Returns:
  117. Dictionary with transcription result
  118. """
  119. # Resample if needed
  120. if sample_rate != 16000:
  121. audio_array = self._resample(audio_array, sample_rate, 16000)
  122. segments, info = self.model.transcribe(
  123. audio_array,
  124. language=self.language,
  125. beam_size=5,
  126. vad_filter=True
  127. )
  128. full_text = ""
  129. segment_list = []
  130. for segment in segments:
  131. text = segment.text.strip()
  132. segment_list.append({
  133. 'start': segment.start,
  134. 'end': segment.end,
  135. 'text': text
  136. })
  137. full_text += text + " "
  138. return {
  139. 'text': full_text.strip(),
  140. 'segments': segment_list,
  141. 'language': info.language,
  142. 'language_probability': info.language_probability
  143. }
  144. def _bytes_to_array(self, audio_bytes: bytes) -> np.ndarray:
  145. """Convert PCM16 bytes to float32 numpy array."""
  146. # Convert bytes to int16
  147. int16_array = np.frombuffer(audio_bytes, dtype=np.int16)
  148. # Convert to float32 in range [-1, 1]
  149. return int16_array.astype(np.float32) / 32768.0
  150. def _resample(self, audio: np.ndarray, orig_sr: int, target_sr: int) -> np.ndarray:
  151. """Simple linear resampling."""
  152. if orig_sr == target_sr:
  153. return audio
  154. duration = len(audio) / orig_sr
  155. new_length = int(duration * target_sr)
  156. indices = np.linspace(0, len(audio) - 1, new_length)
  157. return np.interp(indices, np.arange(len(audio)), audio)
  158. def get_model_info(self) -> dict:
  159. """Get model information."""
  160. return {
  161. 'model_size': self.model_size,
  162. 'device': self.device,
  163. 'compute_type': self.compute_type,
  164. 'language': self.language
  165. }