email-triage-env / inference.py
Janesh's picture
Upload folder using huggingface_hub
73608eb verified
Raw
History Blame Contribute Delete
10 kB
#!/usr/bin/env python3
"""
Email Triage OpenEnv - Baseline Inference Script
Runs a language model agent against all 3 email triage tasks and logs
results in the required OpenEnv [START]/[STEP]/[END] format.
Environment variables:
MODEL_NAME : Model identifier (default: google/gemma-4-31B-it)
API_BASE_URL : LLM API endpoint (default: https://router.huggingface.co/hf-inference/models/google/gemma-4-31B-it/v1)
HF_TOKEN : HuggingFace API token (fallback: OPENAI_API_KEY)
Usage:
export HF_TOKEN="hf_..."
python inference.py
"""
from __future__ import annotations
import logging
import os
import sys
import time
from typing import Dict, List, Optional
from openai import OpenAI
from src.environment import EmailTriageEnv, TASK_CONFIG
from src.grader import grade_task_basic, grade_task_medium, grade_task_hard
from src.models import Action, Observation
from src.utils import format_action_for_log, parse_action_string
# ── Configuration ────────────────────────────────────────────────────────────
API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/hf-inference/models/google/gemma-4-31B-it/v1")
MODEL_NAME = os.getenv("MODEL_NAME", "google/gemma-4-31B-it")
API_KEY = os.getenv("HF_TOKEN") or os.getenv("OPENAI_API_KEY") or ""
TEMPERATURE = 0.3 # low temperature for consistency
MAX_RETRIES = 2
logging.basicConfig(
level=logging.WARNING,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
stream=sys.stderr,
)
logger = logging.getLogger(__name__)
# Task β†’ grader mapping
GRADERS = {
"basic_triage": grade_task_basic,
"multi_folder_triage": grade_task_medium,
"advanced_triage_with_urgency": grade_task_hard,
}
# ── Logging (OpenEnv format β†’ stdout) ────────────────────────────────────────
def log_start(task_id: str, model: str) -> None:
print(f"[START] task={task_id} env=email-triage model={model}", flush=True)
def log_step(
step: int,
action_str: str,
reward: float,
done: bool,
error: Optional[str],
) -> None:
error_part = f'"{error}"' if error else "null"
print(
f"[STEP] step={step} action={action_str} "
f"reward={reward:.2f} done={str(done).lower()} error={error_part}",
flush=True,
)
def log_end(
task_id: str,
success: bool,
steps: int,
score: float,
rewards: List[float],
) -> None:
rewards_str = ",".join(f"{r:.2f}" for r in rewards)
print(
f"[END] task={task_id} success={str(success).lower()} "
f"steps={steps} score={score:.2f} rewards={rewards_str}",
flush=True,
)
# ── Agent ────────────────────────────────────────────────────────────────────
class EmailTriageAgent:
"""Baseline LLM agent for email triage."""
def __init__(self, model_name: str, api_base: str, api_key: str) -> None:
self.client = OpenAI(api_key=api_key, base_url=api_base)
self.model_name = model_name
def get_action(
self,
observation: Observation,
step_num: int,
max_steps: int,
) -> str:
"""Query the LLM to decide which action to take.
Returns the raw action string from the model.
"""
inbox_lines = []
for e in observation.inbox_emails[:10]:
vip_tag = " [VIP]" if e.is_vip_sender else ""
prio_tag = " [URGENT]" if e.priority_flag else ""
inbox_lines.append(
f" ID={e.id} | From: {e.sender} | "
f"Subject: {e.subject[:60]}{vip_tag}{prio_tag}"
)
inbox_text = "\n".join(inbox_lines) if inbox_lines else " (empty)"
prompt = f"""You are an email triage assistant. Sort each email into the correct folder.
INBOX ({len(observation.inbox_emails)} emails):
{inbox_text}
FOLDERS: {', '.join(observation.available_folders)}
RULES:
- Work emails (@acmecorp.com, project updates, code reviews) -> work
- Invoices, expenses, budgets, purchase orders -> finance
- Calendar invites, meeting notes, RSVPs -> meetings
- Phishing, marketing spam, suspicious domains (.ru, .xyz, .biz) -> spam
- Newsletters, notifications, automated reports -> archive
- VIP/urgent emails should be prioritized and classified correctly
Step {step_num}/{max_steps}. Respond with ONLY the action, e.g.: move(0, work)
"""
for attempt in range(MAX_RETRIES + 1):
try:
response = self.client.chat.completions.create(
model=self.model_name,
messages=[{"role": "user", "content": prompt}],
max_tokens=60,
temperature=TEMPERATURE,
timeout=30,
)
return response.choices[0].message.content.strip()
except Exception as e:
logger.warning(
"LLM call failed (attempt %d/%d): %s",
attempt + 1, MAX_RETRIES + 1, e,
)
if attempt < MAX_RETRIES:
time.sleep(2 ** attempt)
# Fallback: move first email to work
if observation.inbox_emails:
return f"move({observation.inbox_emails[0].id}, work)"
return "move(0, work)"
def _parse_or_fallback(
raw_action: str, observation: Observation
) -> Action:
"""Parse the LLM's raw output into an Action, with a safe fallback."""
action = parse_action_string(raw_action)
if action is not None:
return action
# Fallback: move first available email to work
logger.warning("Failed to parse action: '%s'", raw_action)
if observation.inbox_emails:
return Action(
action_type="move",
email_id=observation.inbox_emails[0].id,
target_folder="work",
)
return Action(action_type="move", email_id=0, target_folder="work")
# ── Task Runner ──────────────────────────────────────────────────────────────
def run_task(task_id: str, agent: EmailTriageAgent, seed: int = 42) -> Dict:
"""Run the agent on a single task and log results.
Returns a summary dict with score, steps, success status.
"""
env = EmailTriageEnv(task_id=task_id, seed=seed)
obs = env.reset()
max_steps = env.max_steps
log_start(task_id, MODEL_NAME)
rewards: List[float] = []
steps_taken = 0
error_msg: Optional[str] = None
try:
for step_num in range(1, max_steps + 1):
if obs.done:
break
raw_action = agent.get_action(obs, step_num, max_steps)
action = _parse_or_fallback(raw_action, obs)
action_str = format_action_for_log(action)
try:
obs, reward, done, info = env.step(action)
reward_val = reward.value
step_error = None
except Exception as e:
reward_val = 0.0
done = False
step_error = str(e)
logger.error("Step error: %s", e)
rewards.append(reward_val)
steps_taken = step_num
log_step(
step=step_num,
action_str=action_str,
reward=reward_val,
done=done,
error=step_error,
)
if done:
break
except Exception as e:
error_msg = str(e)
logger.error("Task %s failed: %s", task_id, error_msg)
# Compute final score using the appropriate grader
grader = GRADERS.get(task_id)
if grader and env.history:
score = grader(env.history)
elif rewards:
score = sum(rewards) / len(rewards)
else:
score = 0.0
score = max(0.0, min(1.0, score))
success = score >= TASK_CONFIG[task_id].get("success_threshold", 0.7)
log_end(
task_id=task_id,
success=success,
steps=steps_taken,
score=score,
rewards=rewards,
)
return {
"task_id": task_id,
"score": score,
"steps": steps_taken,
"rewards": rewards,
"success": success,
"error": error_msg,
}
# ── Main ─────────────────────────────────────────────────────────────────────
def main() -> None:
"""Run baseline inference on all 3 tasks."""
if not API_KEY:
logger.error(
"No API key found. Set HF_TOKEN or OPENAI_API_KEY environment variable."
)
sys.exit(1)
agent = EmailTriageAgent(
model_name=MODEL_NAME,
api_base=API_BASE_URL,
api_key=API_KEY,
)
tasks = [
"basic_triage",
"multi_folder_triage",
"advanced_triage_with_urgency",
]
results = []
for task_id in tasks:
result = run_task(task_id, agent, seed=42)
results.append(result)
# Summary to stderr (not stdout, to keep stdout clean for log parsing)
print("\n" + "=" * 60, file=sys.stderr)
print("BASELINE INFERENCE SUMMARY", file=sys.stderr)
print("=" * 60, file=sys.stderr)
for r in results:
status = "PASS" if r["success"] else "FAIL"
print(
f" {r['task_id']:40s} score={r['score']:.2f} [{status}]",
file=sys.stderr,
)
avg = sum(r["score"] for r in results) / len(results) if results else 0.0
print(f"\n Average Score: {avg:.2f}", file=sys.stderr)
print("=" * 60, file=sys.stderr)
if __name__ == "__main__":
main()