joelthomas77 commited on
Commit
fa3a0f5
·
verified ·
1 Parent(s): 60d4850

Make MedASR load failure non-fatal; fall back to Whisper

Browse files
Files changed (1) hide show
  1. app/models/medasr_service.py +30 -16
app/models/medasr_service.py CHANGED
@@ -31,29 +31,38 @@ class MedASRService:
31
  self._load_model()
32
 
33
  def _load_model(self):
34
- """Load MedASR model from Hugging Face."""
 
 
 
 
 
 
 
35
  try:
36
- logger.info(f"Loading MedASR model on device: {self.device}")
37
-
38
- # Use pipeline for easier inference
39
  self.pipe = pipeline(
40
  "automatic-speech-recognition",
41
  model=settings.medasr_model,
42
  device=0 if self.device == "cuda" else -1,
43
  token=settings.hf_token if settings.hf_token else None
44
  )
45
-
46
  logger.info("MedASR model loaded successfully")
47
-
48
- if getattr(settings, "multilingual_asr_enabled", False):
 
 
 
 
 
 
49
  logger.info(f"Loading Whisper model: {settings.whisper_model}")
50
  whisper_name = settings.whisper_model.split("/")[-1].replace("whisper-", "")
51
  self.whisper_model = whisper.load_model(whisper_name, device=self.device)
52
  logger.info("Whisper model loaded successfully")
53
-
54
  except Exception as e:
55
- logger.error(f"Failed to load MedASR model: {e}")
56
- raise
 
57
 
58
  def transcribe(
59
  self,
@@ -111,22 +120,27 @@ class MedASRService:
111
  detected_language = max(probs, key=probs.get)
112
  logger.info(f"Detected language: {detected_language}")
113
 
114
- if detected_language != "en" and self.whisper_model is not None:
 
 
 
 
115
  logger.info(f"Transcribing {detected_language} audio ({duration:.1f}s) with Whisper...")
 
116
  result = self.whisper_model.transcribe(audio_float32, language=detected_language)
117
  transcript = result["text"]
118
  else:
119
  logger.info(f"Transcribing English audio ({duration:.1f}s) with MedASR...")
120
-
121
  # Transcribe using pipeline with chunking for long audio
122
  result = self.pipe(
123
  audio_array,
124
  chunk_length_s=20, # Process in 20-second chunks
125
  stride_length_s=2 # 2-second overlap between chunks
126
  )
127
-
128
  transcript = result["text"]
129
-
130
  # Clean up special tokens and artifacts
131
  import re
132
  transcript = re.sub(r'</?s>|<unk>|<pad>', '', transcript) # Remove special tokens
@@ -142,8 +156,8 @@ class MedASRService:
142
  raise
143
 
144
  def is_ready(self) -> bool:
145
- """Check if the model is loaded and ready."""
146
- return self.pipe is not None
147
 
148
 
149
  # Global instance (singleton pattern)
 
31
  self._load_model()
32
 
33
  def _load_model(self):
34
+ """Load MedASR model from Hugging Face.
35
+
36
+ MedASR (`google/medasr`) is gated and may not be downloadable in
37
+ environments without accepted terms (e.g. free demo deployments).
38
+ We treat its load as best-effort: if it fails, fall back to Whisper
39
+ only and keep the service usable.
40
+ """
41
+ logger.info(f"Loading MedASR model on device: {self.device}")
42
  try:
 
 
 
43
  self.pipe = pipeline(
44
  "automatic-speech-recognition",
45
  model=settings.medasr_model,
46
  device=0 if self.device == "cuda" else -1,
47
  token=settings.hf_token if settings.hf_token else None
48
  )
 
49
  logger.info("MedASR model loaded successfully")
50
+ except Exception as e:
51
+ self.pipe = None
52
+ logger.warning(
53
+ "MedASR model unavailable (%s); falling back to Whisper-only ASR.", e
54
+ )
55
+
56
+ try:
57
+ if getattr(settings, "multilingual_asr_enabled", False) or self.pipe is None:
58
  logger.info(f"Loading Whisper model: {settings.whisper_model}")
59
  whisper_name = settings.whisper_model.split("/")[-1].replace("whisper-", "")
60
  self.whisper_model = whisper.load_model(whisper_name, device=self.device)
61
  logger.info("Whisper model loaded successfully")
 
62
  except Exception as e:
63
+ logger.error("Failed to load Whisper fallback: %s", e)
64
+ if self.pipe is None:
65
+ raise
66
 
67
  def transcribe(
68
  self,
 
120
  detected_language = max(probs, key=probs.get)
121
  logger.info(f"Detected language: {detected_language}")
122
 
123
+ use_whisper = (
124
+ self.pipe is None
125
+ or (detected_language != "en" and self.whisper_model is not None)
126
+ )
127
+ if use_whisper and self.whisper_model is not None:
128
  logger.info(f"Transcribing {detected_language} audio ({duration:.1f}s) with Whisper...")
129
+ audio_float32 = audio_array.astype(np.float32)
130
  result = self.whisper_model.transcribe(audio_float32, language=detected_language)
131
  transcript = result["text"]
132
  else:
133
  logger.info(f"Transcribing English audio ({duration:.1f}s) with MedASR...")
134
+
135
  # Transcribe using pipeline with chunking for long audio
136
  result = self.pipe(
137
  audio_array,
138
  chunk_length_s=20, # Process in 20-second chunks
139
  stride_length_s=2 # 2-second overlap between chunks
140
  )
141
+
142
  transcript = result["text"]
143
+
144
  # Clean up special tokens and artifacts
145
  import re
146
  transcript = re.sub(r'</?s>|<unk>|<pad>', '', transcript) # Remove special tokens
 
156
  raise
157
 
158
  def is_ready(self) -> bool:
159
+ """Check if any ASR backend is available (MedASR pipeline or Whisper fallback)."""
160
+ return self.pipe is not None or self.whisper_model is not None
161
 
162
 
163
  # Global instance (singleton pattern)