grevy
Update.
3e9095a
Raw History Blame Contribute Delete
5.14 kB
import os
from langchain.agents import create_agent
from langchain_core.messages import HumanMessage, SystemMessage
from langchain_groq import ChatGroq
def _make_llm(model: str = "openai/gpt-oss-120b", temperature: float = 0.1) -> ChatGroq:
return ChatGroq(
model=model,
temperature=temperature,
api_key=os.getenv("GROQ_API_KEY"),
timeout=60,
max_retries=2,
)
# Per-role model config. Defaults restricted to models confirmed present on
# this account (openai/gpt-oss-120b and openai/gpt-oss-20b). Override any
# role via env var once you have verified another id with:
# curl -s https://api.groq.com/openai/v1/models \
# -H "Authorization: Bearer $GROQ_API_KEY" | jq '.data[].id'
_ROLE_MODELS = {
"file": os.getenv("GROQ_FILE_MODEL", "openai/gpt-oss-20b"),
"media": os.getenv("GROQ_MEDIA_MODEL", "openai/gpt-oss-120b"),
"normalizer": os.getenv("GROQ_NORMALIZER_MODEL", "openai/gpt-oss-20b"),
}
FILE_ANALYSIS_PROMPT = """You are a document analysis specialist.
You handle questions that reference a task_id, URL, or file path pointing to
a document (PDF, XLSX, CSV, TXT, JSON, MD).
Workflow:
1. If given a task_id (no extension, no slash), call `download_task_file`
to fetch the file and get its local path.
2. Route by extension:
- .pdf -> read_pdf
- .xlsx / .xls -> read_excel
- .csv -> read_csv
- anything else textual -> read_text_file
3. Answer the user's question from the extracted content. Quote values
verbatim when possible. Do not speculate."""
MEDIA_ANALYSIS_PROMPT = """You are a media analysis specialist (image, audio, video).
You handle questions that reference a task_id, URL, or file path pointing to
media, or a YouTube URL.
Workflow:
1. If given a task_id, call `download_task_file` first to get a local path.
2. Route by content type:
- Image (.jpg/.png/.webp/...) -> analyze_image
- Audio (.mp3/.wav/.m4a/.flac/...) -> transcribe_audio, then reason on the text
- Video file or non-YouTube video URL -> analyze_video; if audio matters, also transcribe_audio on the same file
- YouTube URL -> try `youtube_transcript` first; if captions are missing or
the question is about visuals, also call `analyze_video`
3. Give ONLY the answer the question asks for. No preamble."""
GAIA_NORMALIZER_PROMPT = """You are a strict answer normalizer for the GAIA benchmark.
You receive a QUESTION and a RAW answer produced by an agent. Return ONLY the
final answer in exactly the format the question requires:
- Numbers: digits only, no thousand separators, no units unless the question
explicitly asks for units. "3,141" -> "3141". "42 seconds" -> "42" unless
the question asked for a unit.
- Names / entities: just the name, no articles ("the"), no titles, no
explanatory prefix like "The answer is".
- Yes/No: exactly "Yes" or "No".
- Lists: comma-separated, in the order the question asks (alphabetical,
chronological, ...). No trailing "and".
- Dates: match the format the question asks for; if unspecified, use
YYYY-MM-DD.
Never explain, never apologize, never quote the answer.
Output ONLY the normalized answer, on a single line, nothing else."""
class FileAnalysisAgent:
def __init__(self):
from tools import (
download_task_file,
read_csv,
read_excel,
read_pdf,
read_text_file,
)
self.graph = create_agent(
_make_llm(model=_ROLE_MODELS["file"]),
tools=[download_task_file, read_pdf, read_excel, read_csv, read_text_file],
system_prompt=FILE_ANALYSIS_PROMPT,
)
def __call__(self, question: str) -> str:
result = self.graph.invoke(
{"messages": [("user", question)]},
config={"recursion_limit": 8},
)
return result["messages"][-1].content
class MediaAgent:
def __init__(self):
from tools import (
analyze_image,
analyze_video,
download_task_file,
transcribe_audio,
youtube_transcript,
)
self.graph = create_agent(
_make_llm(model=_ROLE_MODELS["media"]),
tools=[
download_task_file,
analyze_image,
analyze_video,
transcribe_audio,
youtube_transcript,
],
system_prompt=MEDIA_ANALYSIS_PROMPT,
)
def __call__(self, question: str) -> str:
result = self.graph.invoke(
{"messages": [("user", question)]},
config={"recursion_limit": 8},
)
return result["messages"][-1].content
class AnswerNormalizer:
def __init__(self):
self.llm = _make_llm(model=_ROLE_MODELS["normalizer"], temperature=0.0)
def __call__(self, question: str, raw_answer: str) -> str:
response = self.llm.invoke([
SystemMessage(content=GAIA_NORMALIZER_PROMPT),
HumanMessage(
content=f"QUESTION:\n{question}\n\nRAW ANSWER:\n{raw_answer}"
),
])
return response.content.strip()