Spaces:
Runtime error
Runtime error
Download inference.py from S777k/openenv-rag-curator: direct link, hf CLI and curl.
- Browser
- Download file 13.9 kB
-
https://huggingface.co/spaces/S777k/openenv-rag-curator/resolve/main/inference.py
- Command line
-
hf download hf://spaces/S777k/openenv-rag-curator/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/S777k/openenv-rag-curator/resolve/main/inference.py
13.9 kB
| """ | |
| inference.py β Baseline inference script for RAG DB Curator OpenEnv | |
| Follows the strict [START] / [STEP] / [END] log format required by the hackathon. | |
| Usage with Hackathon-provided API: | |
| # The hackathon injects these automatically: | |
| # API_BASE_URL - LiteLLM proxy endpoint | |
| # API_KEY - Authentication key | |
| # MODEL_NAME - Model to use | |
| python inference.py | |
| Usage with your own API: | |
| export API_BASE_URL="https://api.openai.com/v1" | |
| export MODEL_NAME="gpt-4o-mini" | |
| export API_KEY="sk-..." | |
| python inference.py | |
| """ | |
| import json | |
| import os | |
| import time | |
| from typing import Any, Dict, List, Optional | |
| import httpx | |
| from openai import OpenAI | |
| # ββ Configuration (from environment variables) βββββββββββββββββββββββββββββββββ | |
| # IMPORTANT: Use API_KEY (not OPENAI_API_KEY or HF_TOKEN) for hackathon compatibility | |
| API_BASE_URL: str = os.environ.get("API_BASE_URL", "https://api.openai.com/v1") | |
| MODEL_NAME: str = os.environ.get("MODEL_NAME", "gpt-4o-mini") | |
| API_KEY: str = os.environ.get("API_KEY", os.environ.get("OPENAI_API_KEY", "")) | |
| if not API_KEY: | |
| print("[ERROR] No API key found. Set API_KEY environment variable.", flush=True) | |
| exit(1) | |
| print(f"[INFO] Using API: {API_BASE_URL}", flush=True) | |
| print(f"[INFO] Model: {MODEL_NAME}", flush=True) | |
| # ββ Environment URL (local or HF Space) ββββββββββββββββββββββββββββββββββββββββ | |
| ENV_BASE_URL: str = os.environ.get("ENV_BASE_URL", "http://localhost:7860") | |
| # ββ Episode settings ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| MAX_STEPS: int = 15 | |
| TEMPERATURE: float = 0.0 | |
| MAX_TOKENS: int = 512 | |
| TASK_IDS: List[str] = ["task_0", "task_1", "task_2"] | |
| # ββ Scoring βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| SUCCESS_SCORE_THRESHOLD: float = 0.5 | |
| BENCHMARK: str = "rag-db-curator" | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Structured log helpers (MANDATORY FORMAT β do not change field names) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def log_start(*, task: str, env: str, model: str) -> None: | |
| """Emit the [START] log line.""" | |
| payload = {"task": task, "env": env, "model": model} | |
| print(f"[START] {json.dumps(payload)}", flush=True) | |
| def log_step(*, | |
| step: int, | |
| action: Any, | |
| reward: float, | |
| done: bool, | |
| error: Optional[str], | |
| ) -> None: | |
| """Emit one [STEP] log line.""" | |
| payload = { | |
| "step": step, | |
| "action": action if isinstance(action, (str, dict)) else str(action), | |
| "reward": round(float(reward), 4), | |
| "done": done, | |
| "error": error, | |
| } | |
| print(f"[STEP] {json.dumps(payload)}", flush=True) | |
| def log_end(*, | |
| success: bool, | |
| steps: int, | |
| score: float, | |
| rewards: List[float], | |
| ) -> None: | |
| """Emit the [END] log line.""" | |
| payload = { | |
| "success": success, | |
| "steps": steps, | |
| "score": round(float(score), 4), | |
| "rewards": [round(float(r), 4) for r in rewards], | |
| } | |
| print(f"[END] {json.dumps(payload)}", flush=True) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Environment HTTP client | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class RAGCuratorClient: | |
| """Thin HTTP wrapper around the FastAPI OpenEnv server.""" | |
| def __init__(self, base_url: str): | |
| self.base_url = base_url.rstrip("/") | |
| self._http = httpx.Client(timeout=30) | |
| def reset(self, task_id: str) -> Dict: | |
| r = self._http.post(f"{self.base_url}/reset/{task_id}") | |
| r.raise_for_status() | |
| return r.json() | |
| def step(self, task_id: str, action: Dict) -> Dict: | |
| r = self._http.post( | |
| f"{self.base_url}/step/{task_id}", | |
| json=action, | |
| headers={"Content-Type": "application/json"}, | |
| ) | |
| r.raise_for_status() | |
| return r.json() | |
| def state(self, task_id: str) -> Dict: | |
| r = self._http.get(f"{self.base_url}/state/{task_id}") | |
| r.raise_for_status() | |
| return r.json() | |
| def close(self) -> None: | |
| self._http.close() | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # LLM helpers | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| SYSTEM_PROMPT = """You are an expert AI Data Engineer working inside a RAG Vector Database | |
| Curator environment. Your job is to clean a messy interview Q&A knowledge base. | |
| You can issue exactly ONE action per turn. Respond with ONLY a valid JSON object | |
| matching one of the schemas below β no markdown fences, no explanation. | |
| Available actions: | |
| {"action_type": "SEARCH_DB", "query": "<keyword>"} | |
| {"action_type": "TAG_QUESTION", "doc_id": "<id>", "tag": "<python|ml|sql|system-design>"} | |
| {"action_type": "UPDATE_ANSWER", "doc_id": "<id>", "answer_text": "<full answer>"} | |
| {"action_type": "MERGE_DUPLICATE", "doc_id": "<keep_id>", "duplicate_doc_id": "<delete_id>"} | |
| {"action_type": "DELETE_DOC", "doc_id": "<id>"} | |
| {"action_type": "SUBMIT_TASK"} | |
| Rules: | |
| - For task_0 (Tagging): Search separately for "python", "ml", "sql", "system-design", "docker", "api" | |
| to find ALL docs with wrong/missing tags. Valid tags: python, ml, sql, system-design. | |
| Also delete junk rows (doc_025, doc_026, doc_027) β they are page boilerplate, not real questions. | |
| - For task_1 (Fill Answers): Search for docs with empty ideal_answer, then UPDATE_ANSWER. | |
| Also delete doc_019 β it has a broken schema (question field = category label). | |
| - For task_2 (Deduplication): Search semantically for similar questions, then MERGE_DUPLICATE. | |
| - When you believe you have completed the task, emit {"action_type": "SUBMIT_TASK"}. | |
| - Never emit anything other than the raw JSON action object.""" | |
| def build_user_message(step: int, obs: Dict, last_reward: float, history: List[str]) -> str: | |
| obs_data = obs.get("observation", obs) | |
| metrics = obs_data.get("database_metrics", {}) | |
| feedback = obs_data.get("feedback", "") | |
| task_desc = obs_data.get("current_task_description", "") | |
| search_hits = obs_data.get("search_results", []) | |
| search_summary = "" | |
| if search_hits: | |
| lines = [] | |
| for hit in search_hits: | |
| lines.append( | |
| f" - {hit['doc_id']}: \"{hit['question']}\" " | |
| f"| answer={'<empty>' if not hit.get('ideal_answer') else hit['ideal_answer'][:60]} " | |
| f"| tags={hit.get('tags', [])}" | |
| ) | |
| search_summary = "Search results:\n" + "\n".join(lines) | |
| recent_history = "\n".join(history[-4:]) if history else "None" | |
| return f"""Step {step} | Last reward: {last_reward:+.3f} | |
| Task: {task_desc} | |
| DB Metrics: {json.dumps(metrics)} | |
| Feedback: {feedback} | |
| {search_summary} | |
| Recent actions: | |
| {recent_history} | |
| What is your next action? Respond with ONLY the JSON action object.""" | |
| def get_model_action(client: OpenAI, step: int, obs: Dict, last_reward: float, history: List[str]) -> Dict: | |
| """Call the LLM and parse its JSON response. Falls back to SUBMIT_TASK on any failure.""" | |
| user_msg = build_user_message(step, obs, last_reward, history) | |
| try: | |
| # Use OpenAI client (compatible with LiteLLM proxy and OpenAI API) | |
| completion = client.chat.completions.create( | |
| model=MODEL_NAME, | |
| messages=[ | |
| {"role": "system", "content": SYSTEM_PROMPT}, | |
| {"role": "user", "content": user_msg}, | |
| ], | |
| temperature=TEMPERATURE, | |
| max_tokens=MAX_TOKENS, | |
| stream=False, | |
| ) | |
| text = (completion.choices[0].message.content or "").strip() | |
| # Strip accidental markdown fences | |
| if text.startswith("```"): | |
| text = text.split("```")[1] | |
| if text.startswith("json"): | |
| text = text[4:] | |
| action = json.loads(text) | |
| if "action_type" not in action: | |
| raise ValueError("Missing action_type") | |
| return action | |
| except Exception as exc: | |
| print(f"[DEBUG] LLM parse error at step {step}: {exc}", flush=True) | |
| return {"action_type": "SUBMIT_TASK"} | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Run one episode for a given task | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def run_episode(client: OpenAI, env: RAGCuratorClient, task_id: str) -> float: | |
| """Runs a full episode for one task. Returns the final normalised score.""" | |
| log_start(task=task_id, env=BENCHMARK, model=MODEL_NAME) | |
| rewards: List[float] = [] | |
| history: List[str] = [] | |
| steps_taken: int = 0 | |
| score: float = 0.0 | |
| success: bool = False | |
| result = env.reset(task_id) | |
| last_reward = 0.0 | |
| try: | |
| for step in range(1, MAX_STEPS + 1): | |
| if result.get("done", False): | |
| break | |
| action = get_model_action(client, step, result, last_reward, history) | |
| error_msg: Optional[str] = None | |
| try: | |
| result = env.step(task_id, action) | |
| except Exception as exc: | |
| error_msg = str(exc) | |
| print(f"[DEBUG] env.step error: {exc}", flush=True) | |
| result = {"done": True, "reward": 0.01, "observation": {}} | |
| reward = float(result.get("reward", 0.01)) | |
| done = result.get("done", False) | |
| last_reward = reward | |
| rewards.append(reward) | |
| steps_taken = step | |
| log_step(step=step, action=action, reward=reward, done=done, error=error_msg) | |
| history.append(f"Step {step}: {action.get('action_type')} β reward {reward:+.3f}") | |
| if done: | |
| score = reward | |
| break | |
| if not rewards: | |
| score = 0.01 # Minimum valid score (strictly between 0 and 1) | |
| elif score == 0.0: | |
| score = rewards[-1] | |
| # Clamp score to be strictly between 0 and 1 (not 0.0, not 1.0) | |
| score = max(0.01, min(score, 0.99)) | |
| success = score >= SUCCESS_SCORE_THRESHOLD | |
| finally: | |
| log_end(success=success, steps=steps_taken, score=score, rewards=rewards) | |
| return score | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Main entry point | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def main() -> None: | |
| # Initialize OpenAI client with hackathon-provided credentials | |
| client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY) | |
| env = RAGCuratorClient(base_url=ENV_BASE_URL) | |
| print(f"[INFO] Model: {MODEL_NAME}", flush=True) | |
| print(f"[INFO] API URL: {API_BASE_URL}", flush=True) | |
| print(f"[INFO] Env URL: {ENV_BASE_URL}", flush=True) | |
| print(f"[INFO] Tasks: {TASK_IDS}", flush=True) | |
| print("", flush=True) | |
| all_scores: Dict[str, float] = {} | |
| try: | |
| for task_id in TASK_IDS: | |
| print(f"{'='*60}", flush=True) | |
| print(f"[INFO] Starting episode for {task_id}", flush=True) | |
| print(f"{'='*60}", flush=True) | |
| t0 = time.time() | |
| score = run_episode(client, env, task_id) | |
| elapsed = time.time() - t0 | |
| all_scores[task_id] = score | |
| print(f"[INFO] {task_id} finished | score={score:.4f} | elapsed={elapsed:.1f}s", flush=True) | |
| print("", flush=True) | |
| finally: | |
| env.close() | |
| print("=" * 60, flush=True) | |
| print("FINAL SCORES", flush=True) | |
| print("=" * 60, flush=True) | |
| for task_id, score in all_scores.items(): | |
| status = "β PASS" if score >= SUCCESS_SCORE_THRESHOLD else "β FAIL" | |
| print(f" {task_id:<12} {score:.4f} {status}", flush=True) | |
| avg = sum(all_scores.values()) / len(all_scores) if all_scores else 0.0 | |
| print(f" {'AVERAGE':<12} {avg:.4f}", flush=True) | |
| print("=" * 60, flush=True) | |
| if __name__ == "__main__": | |
| main() | |