Resolve initial STT language with separate prompts
This commit is contained in:
+121
-22
@@ -284,6 +284,8 @@ _RUSSIAN_LANGUAGE_MARKERS = frozenset({
|
||||
"орыс",
|
||||
"орысша",
|
||||
})
|
||||
_ARABIC_SCRIPT_RE = re.compile(r"[\u0600-\u06ff]")
|
||||
_CYRILLIC_SCRIPT_RE = re.compile(r"[А-Яа-яЁёӘәҒғҚқҢңӨөҰұҮүҺһІі]")
|
||||
|
||||
|
||||
def _detect_preferred_session_language(text: str) -> str | None:
|
||||
@@ -333,6 +335,51 @@ def _detect_session_language(text: str) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _stt_language_candidate_score(text: str, expected_language: str) -> int:
|
||||
normalized_language = expected_language.lower()
|
||||
normalized_text = str(text or "").strip().lower()
|
||||
if not normalized_text:
|
||||
return -100
|
||||
|
||||
score = 0
|
||||
detected_preference = _detect_preferred_session_language(normalized_text)
|
||||
if detected_preference == normalized_language:
|
||||
score += 40
|
||||
elif detected_preference:
|
||||
score -= 40
|
||||
|
||||
if any(ch in _KAZAKH_SPECIFIC_LETTERS for ch in normalized_text):
|
||||
score += 12 if normalized_language == "kk" else -8
|
||||
if any(ch in _RUSSIAN_SPECIFIC_LETTERS for ch in normalized_text):
|
||||
score += 8 if normalized_language == "ru" else -4
|
||||
if _CYRILLIC_SCRIPT_RE.search(normalized_text):
|
||||
score += 4
|
||||
if _ARABIC_SCRIPT_RE.search(normalized_text):
|
||||
score -= 30
|
||||
|
||||
latin_chars = sum(1 for ch in normalized_text if "a" <= ch <= "z")
|
||||
cyrillic_chars = len(_CYRILLIC_SCRIPT_RE.findall(normalized_text))
|
||||
if latin_chars > cyrillic_chars and latin_chars >= 6:
|
||||
score -= 12
|
||||
|
||||
return score
|
||||
|
||||
|
||||
def _choose_language_specific_transcript(ru_transcript: str, kk_transcript: str) -> tuple[str, str | None, int, int]:
|
||||
ru_text = str(ru_transcript or "").strip()
|
||||
kk_text = str(kk_transcript or "").strip()
|
||||
ru_score = _stt_language_candidate_score(ru_text, "ru")
|
||||
kk_score = _stt_language_candidate_score(kk_text, "kk")
|
||||
|
||||
if ru_score >= kk_score and ru_text:
|
||||
detected_language = _detect_name_collection_language(ru_text) or _detect_session_language(ru_text)
|
||||
return ru_text, detected_language, ru_score, kk_score
|
||||
if kk_text:
|
||||
detected_language = _detect_name_collection_language(kk_text) or _detect_session_language(kk_text)
|
||||
return kk_text, detected_language, ru_score, kk_score
|
||||
return ru_text or kk_text, None, ru_score, kk_score
|
||||
|
||||
|
||||
_TTS_LANGUAGE_CODE_MAP = {
|
||||
"ru": "ru",
|
||||
"kk": "kk",
|
||||
@@ -1714,18 +1761,7 @@ class CallSession:
|
||||
len(audio_bytes),
|
||||
_audio_duration_ms(audio_bytes, sample_rate_hz=self.transport.sample_rate_hz),
|
||||
)
|
||||
pending_transcript = (
|
||||
await self._stt.transcribe(
|
||||
audio_bytes,
|
||||
keyterms=self._name_capture_keyterms() if self._awaiting_customer_name else None,
|
||||
language_code=(
|
||||
self._name_capture_language_code()
|
||||
if self._awaiting_customer_name
|
||||
else _stt_language_for(self._session_language)
|
||||
),
|
||||
force_batch=self._awaiting_customer_name,
|
||||
)
|
||||
).strip()
|
||||
pending_transcript = (await self._transcribe_batch_audio(audio_bytes)).strip()
|
||||
except Exception:
|
||||
LOGGER.exception("realtime session %s failed to transcribe pending unanswered audio", self.session_id)
|
||||
continue
|
||||
@@ -1751,6 +1787,78 @@ class CallSession:
|
||||
)
|
||||
return raw or None
|
||||
|
||||
async def _transcribe_batch_audio(self, audio_bytes: bytes) -> str:
|
||||
if self._awaiting_customer_name and self._session_language is None and self._name_capture_language_code() is None:
|
||||
return await self._transcribe_initial_language_choice(audio_bytes)
|
||||
return await self._stt.transcribe(
|
||||
audio_bytes,
|
||||
keyterms=self._name_capture_keyterms() if self._awaiting_customer_name else None,
|
||||
language_code=(
|
||||
self._name_capture_language_code()
|
||||
if self._awaiting_customer_name
|
||||
else _stt_language_for(self._session_language)
|
||||
),
|
||||
force_batch=self._awaiting_customer_name,
|
||||
)
|
||||
|
||||
async def _transcribe_initial_language_choice(self, audio_bytes: bytes) -> str:
|
||||
LOGGER.info(
|
||||
"realtime session %s resolving initial language choice with ru/kk STT requests",
|
||||
self.session_id,
|
||||
)
|
||||
ru_task = self._stt.transcribe(
|
||||
audio_bytes,
|
||||
keyterms=self._name_capture_keyterms(),
|
||||
language_code="ru",
|
||||
force_batch=True,
|
||||
)
|
||||
kk_task = self._stt.transcribe(
|
||||
audio_bytes,
|
||||
keyterms=self._name_capture_keyterms(),
|
||||
language_code="kk",
|
||||
force_batch=True,
|
||||
)
|
||||
ru_raw, kk_raw = await asyncio.gather(ru_task, kk_task, return_exceptions=True)
|
||||
if isinstance(ru_raw, Exception):
|
||||
LOGGER.warning(
|
||||
"realtime session %s initial Russian STT request failed: %s",
|
||||
self.session_id,
|
||||
ru_raw,
|
||||
)
|
||||
ru_result = ""
|
||||
else:
|
||||
ru_result = str(ru_raw or "")
|
||||
if isinstance(kk_raw, Exception):
|
||||
LOGGER.warning(
|
||||
"realtime session %s initial Kazakh STT request failed: %s",
|
||||
self.session_id,
|
||||
kk_raw,
|
||||
)
|
||||
kk_result = ""
|
||||
else:
|
||||
kk_result = str(kk_raw or "")
|
||||
transcript, detected_language, ru_score, kk_score = _choose_language_specific_transcript(ru_result, kk_result)
|
||||
LOGGER.info(
|
||||
"realtime session %s initial language STT selected: detected=%s ru_score=%s kk_score=%s "
|
||||
"ru=%r kk=%r selected=%r",
|
||||
self.session_id,
|
||||
detected_language,
|
||||
ru_score,
|
||||
kk_score,
|
||||
_preview_text(ru_result),
|
||||
_preview_text(kk_result),
|
||||
_preview_text(transcript),
|
||||
)
|
||||
if detected_language and self._session_language is None:
|
||||
self._session_language = detected_language
|
||||
LOGGER.info(
|
||||
"realtime session %s language locked during initial language choice: language=%s transcript=%r",
|
||||
self.session_id,
|
||||
detected_language,
|
||||
_preview_text(transcript),
|
||||
)
|
||||
return transcript
|
||||
|
||||
def _build_name_collection_response(self, transcript: str) -> str | None:
|
||||
if not self._awaiting_customer_name:
|
||||
return None
|
||||
@@ -2386,16 +2494,7 @@ class CallSession:
|
||||
len(utterance_audio),
|
||||
_audio_duration_ms(utterance_audio, sample_rate_hz=self.transport.sample_rate_hz),
|
||||
)
|
||||
return await self._stt.transcribe(
|
||||
utterance_audio,
|
||||
keyterms=self._name_capture_keyterms() if self._awaiting_customer_name else None,
|
||||
language_code=(
|
||||
self._name_capture_language_code()
|
||||
if self._awaiting_customer_name
|
||||
else _stt_language_for(self._session_language)
|
||||
),
|
||||
force_batch=self._awaiting_customer_name,
|
||||
)
|
||||
return await self._transcribe_batch_audio(utterance_audio)
|
||||
|
||||
def _maybe_dump_audio(self, audio_bytes: bytes) -> None:
|
||||
if str(os.getenv("ENABLE_AUDIO_DUMP", "")).strip().lower() not in {"1", "true", "yes", "on"}:
|
||||
|
||||
Reference in New Issue
Block a user