OutstandingOm/knowledge-graph-env
1
1#!/usr/bin/env python32"""3Inference script for Knowledge Graph Environment.4 5Runs 3 tasks (task_easy, task_medium, task_hard), grades each via the6/grade HTTP endpoint, and emits the required stdout format:7 8 [START] task=<task_name> env=<benchmark> model=<model_name>9 [STEP] step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>10 [END] success=<true|false> steps=<n> rewards=<r1,r2,...,rn>11"""12 13import os14import sys15import asyncio16from typing import List, Optional17 18import httpx19from openai import OpenAI20 21# ── Environment variables ─────────────────────────────────────────────────────22HF_TOKEN = os.getenv("HF_TOKEN")23if HF_TOKEN is None:24 raise ValueError("HF_TOKEN environment variable is required")25 26API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")27MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")28API_KEY = HF_TOKEN29# Space URL — during hackathon eval this is the live HF Space30ENV_URL = os.getenv("ENV_URL",31 "https://outstandingom-knowledge-graph-env.hf.space").rstrip("/")32BENCHMARK = "knowledge_graph_env"33MAX_STEPS = 334SUCCESS_THRESHOLD = 0.135 36# ── Task definitions ──────────────────────────────────────────────────────────37# Each task has an initial input and a system prompt for the LLM agent.38TASK_SPECS = {39 "task_easy": {40 "input": "I cannot login to my account. My password is not working and I keep getting locked out.",41 "system": (42 "You are a customer support AI. Identify the core issue the user is facing. "43 "Respond with a concise identification of the problem (1-2 sentences)."44 ),45 },46 "task_medium": {47 "input": "My bill shows a double charge for my subscription this month. I need a refund for the extra payment Invoice INV-2024-891.",48 "system": (49 "You are a customer support AI specialising in billing. Identify the billing issue "50 "and what action should be taken. Respond concisely (1-2 sentences)."51 ),52 },53 "task_hard": {54 "input": "My account is locked after multiple failed password attempts. I suspect a security breach. Please help urgently — critical issue.",55 "system": (56 "You are a security-aware customer support AI. Identify the critical security issue "57 "and recommend the immediate action. Respond concisely (1-2 sentences)."58 ),59 },60}61 62# ── Logging helpers ───────────────────────────────────────────────────────────63def log_start(task: str, env: str, model: str) -> None:64 print(f"[START] task={task} env={env} model={model}", flush=True)65 66def log_step(step: int, action: str, reward: float, done: bool,67 error: Optional[str]) -> None:68 safe_action = action.replace("\n", " ").replace("\r", "")[:120]69 print(70 f"[STEP] step={step} action={safe_action} "71 f"reward={reward:.2f} done={str(done).lower()} "72 f"error={error if error else 'null'}",73 flush=True,74 )75 76def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:77 # Notice we keep 'score' in the function arguments to calculate success, 78 # but we DO NOT print it in the string below to comply with the Hackathon Regex parser.79 r_str = ",".join(f"{r:.2f}" for r in rewards)80 print(81 f"[END] success={str(success).lower()} steps={steps} rewards={r_str}",82 flush=True,83 )84 85# ── LLM helper ────────────────────────────────────────────────────────────────86def get_llm_action(client: OpenAI, system: str, user: str) -> str:87 """Ask the LLM for an action; falls back to the raw input on any error."""88 try:89 resp = client.chat.completions.create(90 model=MODEL_NAME,91 messages=[92 {"role": "system", "content": system},93 {"role": "user", "content": user},94 ],95 temperature=0.3,96 max_tokens=120,97 )98 text = (resp.choices[0].message.content or "").strip()99 return text if text else user100 except Exception as exc:101 print(f"[DEBUG] LLM call failed: {exc}", file=sys.stderr, flush=True)102 return user103 104# ── Grade helper ──────────────────────────────────────────────────────────────105async def grade(http: httpx.AsyncClient, task_id: str, action: str) -> float:106 """POST to /grade and return the score, or 0.01 on failure."""107 try:108 r = await http.post(109 f"{ENV_URL}/grade",110 json={"task_id": task_id, "input_text": action},111 timeout=30.0,112 )113 r.raise_for_status()114 return float(r.json().get("score", 0.01))115 except Exception as exc:116 print(f"[DEBUG] /grade failed for {task_id}: {exc}", file=sys.stderr, flush=True)117 return 0.01118 119# ── Single-task runner ────────────────────────────────────────────────────────120async def run_task(121 task_id: str,122 spec: dict,123 llm_client: OpenAI,124 http: httpx.AsyncClient,125) -> None:126 log_start(task=task_id, env=BENCHMARK, model=MODEL_NAME)127 128 rewards: List[float] = []129 steps_done = 0130 success = False131 score = 0.0132 observation = spec["input"]133 134 for step in range(1, MAX_STEPS + 1):135 # Agent generates action from LLM136 action = get_llm_action(llm_client, spec["system"], observation)137 138 # Grade the action via the environment's /grade endpoint139 reward = await grade(http, task_id, action)140 reward = max(0.01, min(0.99, reward)) # strict (0, 1) range141 142 done = (step == MAX_STEPS)143 error = None144 145 rewards.append(reward)146 steps_done = step147 148 log_step(step=step, action=action, reward=reward, done=done, error=error)149 150 if done:151 break152 153 # For subsequent steps, feed the graded score back as context154 observation = (155 f"Previous action: {action}\n"156 f"Reward received: {reward:.4f}\n"157 f"Original issue: {spec['input']}"158 )159 160 score = sum(rewards) / len(rewards) if rewards else 0.0161 score = max(0.0, min(1.0, score))162 success = score >= SUCCESS_THRESHOLD163 164 log_end(success=success, steps=steps_done, score=score, rewards=rewards)165 166# ── Main ──────────────────────────────────────────────────────────────────────167async def main() -> None:168 llm_client = OpenAI(169 base_url=API_BASE_URL,170 api_key=API_KEY if API_KEY else "no-key",171 )172 173 async with httpx.AsyncClient() as http:174 for task_id, spec in TASK_SPECS.items():175 await run_task(task_id, spec, llm_client, http)176 177if __name__ == "__main__":178 asyncio.run(main())179 