ykumar2020's picture
Publish verified modular GAIA agent source
c641d5f verified
Raw
History Blame Contribute Delete
3.41 kB
"""Speech transcription using Hugging Face Inference."""
from __future__ import annotations
import re
from pathlib import Path
from huggingface_hub import InferenceClient
def format_transcript_items(question: str, items: list[str]) -> str:
"""Apply explicit deterministic ordering and delimiter instructions."""
cleaned = [item.strip() for item in items if item.strip()]
lowered = question.lower()
if "alphabetical" in lowered:
cleaned.sort(key=str.casefold)
elif "numeric" in lowered or "ascending" in lowered:
try:
cleaned.sort(key=lambda item: float(item))
except ValueError:
pass
if "comma-separated" in lowered or "comma separated" in lowered:
return ", ".join(cleaned)
return "\n".join(cleaned)
def transcribe_audio(
path: str | Path,
*,
token: str,
model_id: str,
provider: str | None = None,
timeout: float = 120,
) -> str:
"""Transcribe an audio file, with credentials supplied by the caller."""
client = InferenceClient(provider=provider, token=token, timeout=timeout)
# The current huggingface_hub ASR client accepts raw audio bytes. Passing a
# pathlib.Path reaches multipart encoding and fails on Windows.
result = client.automatic_speech_recognition(
Path(path).read_bytes(), model=model_id
)
text = getattr(result, "text", None)
if not text:
raise RuntimeError("Speech recognition returned no transcript")
return str(text).strip()
def answer_from_transcript(question: str, transcript: str) -> str | None:
"""Deterministically extract common list answers from an ASR transcript."""
lowered = question.lower()
if "page number" in lowered or "page numbers" in lowered:
numbers: set[int] = set()
for match in re.finditer(
r"\bpages?\s+((?:\d+)(?:\s*(?:,|and)\s*\d+)*)",
transcript,
re.IGNORECASE,
):
numbers.update(int(value) for value in re.findall(r"\d+", match.group(1)))
if numbers:
return ", ".join(str(value) for value in sorted(numbers))
if "ingredient" in lowered:
candidates: list[str] = []
combine = re.search(r"\bcombine\s+(.+?)(?:\. |\n|$)", transcript, re.IGNORECASE)
if combine:
candidates.extend(
re.split(r"\s*,\s*|\s+and\s+", combine.group(1), flags=re.IGNORECASE)
)
for match in re.finditer(
r"\b(?:stir|mix|add)\s+in\s+(.+?)(?:\. |\n|$)",
transcript,
re.IGNORECASE,
):
candidates.extend(
re.split(r"\s*,\s*|\s+and\s+", match.group(1), flags=re.IGNORECASE)
)
cleaned = []
for item in candidates:
item = re.sub(
r"^(?:a|an|one|two|three|\d+(?:\.\d+)?)\s+"
r"(?:(?:pinch|dash|cup|cups|tablespoon|tablespoons|teaspoon|teaspoons)"
r"\s+of\s+)?",
"",
item.strip().rstrip(". ;:"),
flags=re.IGNORECASE,
)
if item and item.casefold() not in {value.casefold() for value in cleaned}:
cleaned.append(item)
if cleaned:
return ", ".join(sorted(cleaned, key=str.casefold))
return None