Team Ai
Apppublic

training-monkey/dataoncallenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py310 linesDownload Raw Back to root
1"""Baseline inference script for DataOnCallEnv.2"""3 4import os5import json6import time7import httpx8from openai import OpenAI9from environment import DataOnCallEnv10from models import Action11from dotenv import load_dotenv12 13load_dotenv()14 15# Configuration16API_BASE_URL = os.environ.get("API_BASE_URL", "https://router.huggingface.co/v1")17MODEL_NAME   = os.environ.get("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct").strip()18 19API_KEY = (20    os.environ.get("HF_TOKEN") or21    os.environ.get("OPENAI_API_KEY")22)23 24if not API_KEY:25    raise EnvironmentError("Set HF_TOKEN or OPENAI_API_KEY before running.")26 27client = OpenAI(api_key=API_KEY, base_url=API_BASE_URL)28 29SUCCESS_SCORE_THRESHOLD = 0.5  # Treat 0.5+ as success for the logger30 31# Validator logger functions32def log_start(task: str, env: str, model: str) -> None:33    print(f"[START] task={task} env={env} model={model}", flush=True)34 35def log_step(step: int, action: str, reward: float, done: bool, error: str) -> None:36    error_val = error if error else "null"37    done_val = str(done).lower()38    # Remove newlines from action string just in case39    action_clean = action.replace("\n", " ").replace("\r", "")40    print(41        f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}",42        flush=True,43    )44 45def log_end(success: bool, steps: int, score: float, rewards: list) -> None:46    rewards_str = ",".join(f"{r:.2f}" for r in rewards)47    print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)48 49 50# Tool definitions51TOOLS = [52    {53        "type": "function",54        "function": {55            "name": "run_sql",56            "description": (57                "Execute a SELECT SQL query. "58                "Specify column names — SELECT * is blocked. "59                "Discover tables first with list_tables(). Cost: 2.0"60            ),61            "parameters": {62                "type": "object",63                "properties": {64                    "query":     {"type": "string", "description": "The SELECT query"},65                    "reasoning": {"type": "string", "description": "Why you are running this"},66                },67                "required": ["query", "reasoning"],68            },69        },70    },71    {72        "type": "function",73        "function": {74            "name": "inspect_schema",75            "description": "See column names/types. Discover tables first. Cost: 1.0",76            "parameters": {77                "type": "object",78                "properties": {79                    "query":     {"type": "string", "description": "Table name"},80                    "reasoning": {"type": "string", "description": "Why inspecting"},81                },82                "required": ["query", "reasoning"],83            },84        },85    },86    {87        "type": "function",88        "function": {89            "name": "check_logs",90            "description": "Read the dbt changelog. Cost: 1.0",91            "parameters": {92                "type": "object",93                "properties": {94                    "query":     {"type": "string", "description": "Pass empty str"},95                    "reasoning": {"type": "string"},96                },97                "required": ["query", "reasoning"],98            },99        },100    },101    {102        "type": "function",103        "function": {104            "name": "check_airflow",105            "description": "Read Airflow run history. Cost: 1.0",106            "parameters": {107                "type": "object",108                "properties": {109                    "query":     {"type": "string", "description": "Pass empty str"},110                    "reasoning": {"type": "string"},111                },112                "required": ["query", "reasoning"],113            },114        },115    },116    {117        "type": "function",118        "function": {119            "name": "diff_report",120            "description": "Compare dates. Format: 'YYYY-MM-DD,YYYY-MM-DD'. Cost: 1.5",121            "parameters": {122                "type": "object",123                "properties": {124                    "query":     {"type": "string", "description": "'date1,date2'"},125                    "reasoning": {"type": "string"},126                },127                "required": ["query", "reasoning"],128            },129        },130    },131    {132        "type": "function",133        "function": {134            "name": "list_tables",135            "description": "Discover tables. MUST call this first. Cost: 0.5",136            "parameters": {137                "type": "object",138                "properties": {139                    "query":     {"type": "string", "description": "Pass empty str"},140                    "reasoning": {"type": "string"},141                },142                "required": ["query", "reasoning"],143            },144        },145    },146    {147        "type": "function",148        "function": {149            "name": "submit",150            "description": "End episode. Supply root cause + SQL. Cost: 0.0",151            "parameters": {152                "type": "object",153                "properties": {154                    "query": {"type": "string", "description": "ROOT CAUSE: ... CORRECTED SQL: ..."},155                    "reasoning": {"type": "string"},156                },157                "required": ["query", "reasoning"],158            },159        },160    },161]162 163SYSTEM_PROMPT = """You are a Data Analyst debugging reports.164RULES:165- Budget: 15 steps, 20.0 cost. Tables HIDDEN until list_tables() used.166- SELECT * is BLOCKED. Specify columns. min 2 steps before submit.167 168STRATEGY:1691. list_tables() and inspect_schema() first.1702. check_logs()/check_airflow() for root causes.1713. run_sql() to confirm hypothesis.1724. submit() immediately when diagnosed.173 174SUBMIT FORMAT:175ROOT CAUSE: [explanation]176CORRECTED SQL: [full SELECT query fixing the bug]"""177 178# Agent loop179 180def run_agent(env: DataOnCallEnv, task_id: int):181    task_name = ["", "weekly_revenue", "monthly_active_users", "revenue_inflation"][task_id]182    183    log_start(task=task_name, env="DataOnCallEnv", model=MODEL_NAME)184    185    obs = env.reset(task_id=task_id)186 187    messages = [188        {"role": "system", "content": SYSTEM_PROMPT},189        {"role": "user",   "content": json.dumps(obs.result, indent=2)},190    ]191 192    final_reward = None193    step_rewards = []194    steps        = 0195    MAX_STEPS    = env.MAX_STEPS196 197    while not env.done and steps < MAX_STEPS:198 199        steps_remaining = MAX_STEPS - steps200        budget_remaining = env.COST_BUDGET - env.cost_spent201        202        if steps_remaining == 5 or budget_remaining <= 5.0:203            messages.append({204                "role": "user",205                "content": f"WARNING: {steps_remaining} steps, budget: {budget_remaining:.1f}. Call submit()NOW if ready."206            })207        elif steps_remaining == 2:208            messages.append({209                "role": "user",210                "content": "FINAL WARNING: 2 steps left. Call submit() next step."211            })212 213        max_retries = 5214        for attempt in range(max_retries):215            try:216                response = client.chat.completions.create(217                    model=MODEL_NAME,218                    messages=messages,219                    tools=TOOLS,220                    tool_choice="auto",221                    temperature=0,222                )223                break224            except Exception as e:225                err_msg = str(e).lower()226                if "503" in err_msg or "loading" in err_msg:227                    time.sleep(15 * (attempt + 1))228                    if attempt == max_retries - 1: raise e229                elif "429" in err_msg:230                    time.sleep(30)231                    if attempt == max_retries - 1: raise e232                else:233                    raise e234 235        msg = response.choices[0].message236 237        if msg.tool_calls:238            messages.append(msg)239            for tool_call in msg.tool_calls:240                fn_name = tool_call.function.name241                try:242                    fn_args = json.loads(tool_call.function.arguments)243                except Exception:244                    fn_args = {}245 246                query     = fn_args.get("query", "")247                reasoning = fn_args.get("reasoning", "")248                249                action = Action(tool=fn_name, query=query, reasoning=reasoning)250                obs, reward, done, info = env.step(action)251                steps += 1252                253                # Keep episode running rewards tracked properly254                current_reward_val = reward.score if reward else 0.0255                step_rewards.append(current_reward_val)256                257                if reward: final_reward = reward258 259                # Derive error state for logging260                step_error = None261                if isinstance(obs.result, dict):262                    if obs.result.get("error") or not obs.result.get("success", True):263                        step_error = str(obs.result.get("error") or obs.result.get("result", "error"))264                        step_error = step_error.replace("\n", " ").replace("\r", "")265 266                action_str = f"{fn_name}({query})"267                log_step(steps, action_str, current_reward_val, done, step_error)268 269                messages.append({270                    "role":         "tool",271                    "tool_call_id": tool_call.id,272                    "content":      json.dumps(obs.result),273                })274 275                if done:276                    break277        else:278            steps += 1279            content = msg.content or "No Tool Used"280            log_step(steps, "invalid_output", 0.0, False, "Model did not output a tool call")281            282            step_rewards.append(0.0)283            messages.append({"role": "assistant", "content": content})284            messages.append({285                "role":    "user",286                "content": "Use a tool to continue, or call submit() with your answer.",287            })288 289    # Backup grade calculation if ended abruptly290    if final_reward is None:291        from graders import grade292        final_reward = grade(env.task_id, env.conn, env.actions, env.final_answer)293        if len(step_rewards) == 0:294            step_rewards.append(final_reward.score)295        else:296            step_rewards[-1] = final_reward.score297            298    is_success = final_reward.score >= SUCCESS_SCORE_THRESHOLD299    log_end(is_success, steps, final_reward.score, step_rewards)300    301    return {302        "task_id": task_id,303        "score": final_reward.score,304        "steps_taken": steps305    }306 307if __name__ == "__main__":308    env = DataOnCallEnv()309    for task_id in [1, 2, 3]:310        run_agent(env, task_id)