gojiteji's picture
Audio Decision Model demo
148af80
Raw History Blame Contribute Delete
5.12 kB
"""Schema-gated language questions answered by Qwen3-ASR's native tag."""
from __future__ import annotations
from dataclasses import dataclass
import re
from typing import Any, Mapping
from .typed_audio import question_options, typed_answer
LANGUAGE_ALIASES: dict[str, tuple[str, ...]] = {
"Chinese": ("zh", "chinese", "mandarin", "中国語", "中文", "普通話"),
"English": ("en", "english", "英語"),
"Cantonese": ("yue", "cantonese", "広東語", "粤語"),
"Arabic": ("ar", "arabic", "アラビア語"),
"German": ("de", "german", "ドイツ語"),
"French": ("fr", "french", "フランス語"),
"Spanish": ("es", "spanish", "スペイン語"),
"Portuguese": ("pt", "portuguese", "ポルトガル語"),
"Indonesian": ("id", "indonesian", "インドネシア語"),
"Italian": ("it", "italian", "イタリア語"),
"Korean": ("ko", "korean", "韓国語", "朝鮮語"),
"Russian": ("ru", "russian", "ロシア語"),
"Thai": ("th", "thai", "タイ語"),
"Vietnamese": ("vi", "vietnamese", "ベトナム語"),
"Japanese": ("ja", "jp", "japanese", "日本語"),
"Turkish": ("tr", "turkish", "トルコ語"),
"Hindi": ("hi", "hindi", "ヒンディー語", "ヒンディ語"),
"Malay": ("ms", "malay", "マレー語"),
"Dutch": ("nl", "dutch", "オランダ語"),
"Swedish": ("sv", "swedish", "スウェーデン語"),
"Danish": ("da", "danish", "デンマーク語"),
"Finnish": ("fi", "finnish", "フィンランド語"),
"Polish": ("pl", "polish", "ポーランド語"),
"Czech": ("cs", "czech", "チェコ語"),
"Filipino": ("fil", "filipino", "フィリピン語", "タガログ語"),
"Persian": ("fa", "persian", "ペルシャ語"),
"Greek": ("el", "greek", "ギリシャ語"),
"Romanian": ("ro", "romanian", "ルーマニア語"),
"Hungarian": ("hu", "hungarian", "ハンガリー語"),
"Macedonian": ("mk", "macedonian", "マケドニア語"),
}
LANGUAGE_CUES = ("language", "spoken", "speaking", "言語", "何語", "話され", "話して", "発話")
def _contains(text: str, alias: str) -> bool:
value = re.sub(r"\s+", " ", text.casefold()).strip()
needle = re.sub(r"\s+", " ", alias.casefold()).strip()
if needle.isascii() and any(char.isalpha() for char in needle):
return re.search(r"(?<![a-z])" + re.escape(needle) + r"(?![a-z])", value) is not None
return needle in value
def _languages(text: str) -> tuple[str, ...]:
return tuple(language for language, aliases in LANGUAGE_ALIASES.items()
if any(_contains(text, alias) for alias in aliases))
@dataclass(frozen=True)
class LanguageQuestion:
task: str
option_languages: tuple[str, ...] = ()
queried_language: str | None = None
def classify_language_question(question: Mapping[str, Any]) -> LanguageQuestion | None:
"""Recognize explicit language choice/presence schemas without using question IDs."""
pairs = question_options(dict(question))
if question.get("type") == "choice":
resolved = []
for key, description in pairs:
candidates = _languages(f"{key}\n{description}")
if len(candidates) != 1:
return None
resolved.append(candidates[0])
if len(set(resolved)) != len(resolved):
return None
instruction = str(question.get("instructions", ""))
if not any(_contains(instruction, cue) for cue in LANGUAGE_CUES) and not _languages(instruction):
# Canonical language IDs in every option are sufficient even when the
# instruction is simply "どれですか".
if not all(_languages(key) for key, _ in pairs):
return None
return LanguageQuestion("language_id", tuple(resolved))
if question.get("type") == "noul":
candidates = _languages(str(question.get("instructions", "")))
if len(candidates) == 1:
return LanguageQuestion("language_presence", queried_language=candidates[0])
return None
def answer_language_questions(detected: str, questions: Mapping[str, Mapping[str, Any]]) -> dict[str, dict]:
"""Map one or more canonical Qwen language tags into typed answers."""
observed = {value.strip() for value in detected.split(",") if value.strip()}
answers: dict[str, dict] = {}
for name, question in questions.items():
task = classify_language_question(question)
if task is None:
raise ValueError("Unsupported language question")
if task.task == "language_presence":
present = task.queried_language in observed
answers[name] = typed_answer(dict(question), [0.0, 1.0] if present else [1.0, 0.0])
else:
probabilities = [1.0 if language in observed else 0.0 for language in task.option_languages]
if not any(probabilities):
probabilities = [1.0] * len(probabilities)
answers[name] = typed_answer(dict(question), probabilities)
answers[name]["confidence_kind"] = "qwen_asr_language_tag_not_calibrated"
return answers