Team Ai
Apppublic

violinadoley25/multi-agent-task-alloc

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py241 linesDownload Raw Back to root
1"""2inference.py - Multi-Agent Task Allocation inference script.3 4Runs all 3 tasks (easy_allocation, medium_allocation, hard_allocation) sequentially5using an OpenAI-compatible LLM to generate allocation decisions.6 7Environment variables8---------------------9API_BASE_URL   OpenAI-compatible API base URL10               (default: https://router.huggingface.co/v1)11MODEL_NAME     Model identifier12               (default: Qwen/Qwen2.5-72B-Instruct)13HF_TOKEN       HuggingFace / API token (required for default endpoint)14IMAGE_NAME     Docker image name (default: multi-agent-task-alloc:latest)15ENV_URL        Override to connect to a running server instead of Docker16 17Stdout format18-------------19[START] task=<task_name> env=<benchmark> model=<model_name>20[STEP]  step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>21[END]   success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...,rn>22"""23 24from __future__ import annotations25 26import asyncio27import json28import os29import sys30from typing import Any, Dict, List, Optional31 32from openai import OpenAI33 34from multi_agent_task_alloc import TaskAllocationEnv, TaskAllocationAction35 36# ─── Configuration ────────────────────────────────────────────────────────────37 38API_BASE_URL: str = os.getenv("API_BASE_URL") or "https://router.huggingface.co/v1"39MODEL_NAME: str = os.getenv("MODEL_NAME") or "Qwen/Qwen2.5-72B-Instruct"40API_KEY: str = os.getenv("HF_TOKEN") or os.getenv("API_KEY") or "no-key"41IMAGE_NAME: str = os.getenv("IMAGE_NAME") or "multi-agent-task-alloc:latest"42ENV_URL: Optional[str] = os.getenv("ENV_URL")43 44BENCHMARK = "multi-agent-task-alloc"45MAX_STEPS = 1546TEMPERATURE = 0.147MAX_TOKENS = 25648SUCCESS_THRESHOLD = 0.549 50TASKS = [51    {"index": 0, "name": "easy_allocation"},52    {"index": 1, "name": "medium_allocation"},53    {"index": 2, "name": "hard_allocation"},54]55 56# ─── OpenAI client ────────────────────────────────────────────────────────────57 58_llm: Optional[OpenAI] = None59 60 61def get_llm() -> OpenAI:62    global _llm63    if _llm is None:64        _llm = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)65    return _llm66 67 68# ─── Logging ──────────────────────────────────────────────────────────────────69 70def log_start(task: str, env: str, model: str) -> None:71    print(f"[START] task={task} env={env} model={model}", flush=True)72 73 74def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:75    safe = action.replace("\n", " ").replace("\r", "")[:120]76    err = error if error else "null"77    done_str = "true" if done else "false"78    print(f"[STEP] step={step} action={safe} reward={reward:.2f} done={done_str} error={err}", flush=True)79 80 81def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:82    rewards_str = ",".join(f"{r:.2f}" for r in rewards) if rewards else "0.00"83    success_str = "true" if success else "false"84    print(f"[END] success={success_str} steps={steps} score={score:.2f} rewards={rewards_str}", flush=True)85 86 87# ─── Prompt builder ───────────────────────────────────────────────────────────88 89SYSTEM_PROMPT = (90    "You are a project manager AI. Assign tasks to team members based on their skills.\n"91    "Rules:\n"92    "- Match agent skills to task required_skills for best results.\n"93    "- Don't overload agents (respect their load capacity).\n"94    "- Respond with ONLY valid JSON: {\"task_id\": \"<id>\", \"agent_id\": \"<id>\"}\n"95    "- No explanation, no markdown, just the JSON object."96)97 98 99def _build_prompt(obs: Any) -> str:100    state = obs.game_state101    pending = state.get("pending_tasks", [])102    team = state.get("team", [])103 104    tasks_str = "\n".join(105        f"  - {t['id']}: {t['name']} (requires: {', '.join(t['required_skills']) or 'any'})"106        for t in pending107    )108    team_str = "\n".join(109        f"  - {a['id']}: {a['name']} | skills: {', '.join(a['skills'])} | load: {a['load']}"110        for a in team111    )112 113    return (114        f"Pending tasks:\n{tasks_str}\n\n"115        f"Available team members:\n{team_str}\n\n"116        f"Pick ONE task and ONE agent. Return JSON only."117    )118 119 120def generate_action(obs: Any) -> str:121    prompt = _build_prompt(obs)122    completion = get_llm().chat.completions.create(123        model=MODEL_NAME,124        messages=[125            {"role": "system", "content": SYSTEM_PROMPT},126            {"role": "user", "content": prompt},127        ],128        temperature=TEMPERATURE,129        max_tokens=MAX_TOKENS,130    )131    return (completion.choices[0].message.content or "").strip()132 133 134# ─── Task runner ─────────────────────────────────────────────────────────────135 136async def run_task(137    env: TaskAllocationEnv,138    task_index: int,139    task_name: str,140) -> Dict[str, Any]:141    log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)142 143    rewards: List[float] = []144    steps_taken = 0145    final_score = 0.0146    success = False147    last_error: Optional[str] = None148 149    try:150        result = await env.reset(task_index=task_index)151        obs = result.observation152 153        for step_n in range(1, MAX_STEPS + 1):154            if obs.done:155                break156 157            # Stop if no pending tasks158            if not obs.game_state.get("pending_tasks"):159                break160 161            try:162                content = generate_action(obs)163            except Exception as exc:164                last_error = f"LLM error: {exc}"165                log_step(step=step_n, action="null", reward=0.0, done=False, error=last_error)166                break167 168            try:169                result = await env.step(TaskAllocationAction(action_type="allocate", content=content))170            except Exception as exc:171                last_error = f"Env error: {exc}"172                log_step(step=step_n, action=content, reward=0.0, done=False, error=last_error)173                break174 175            obs = result.observation176            reward = float(result.reward) if result.reward is not None else 0.0177            done = bool(result.done)178            steps_taken = step_n179            rewards.append(reward)180            last_error = None181 182            log_step(step=step_n, action=content, reward=reward, done=done, error=None)183 184            if done:185                final_score = float(obs.score) if obs.score is not None else 0.0186                final_score = max(0.0, min(1.0, final_score))187                success = final_score >= SUCCESS_THRESHOLD188                break189 190    except Exception as exc:191        last_error = str(exc)192        print(f"[DEBUG] run_task error: {exc}", file=sys.stderr, flush=True)193 194    finally:195        log_end(success=success, steps=steps_taken, score=final_score, rewards=rewards)196 197    return {198        "task_name": task_name,199        "success": success,200        "steps": steps_taken,201        "score": final_score,202        "rewards": rewards,203        "error": last_error,204    }205 206 207# ─── Main ─────────────────────────────────────────────────────────────────────208 209async def main() -> None:210    if ENV_URL:211        print(f"Connecting to environment at {ENV_URL} ...", flush=True)212        env = TaskAllocationEnv(base_url=ENV_URL)213    else:214        print(f"Starting environment from Docker image: {IMAGE_NAME} ...", flush=True)215        env = await TaskAllocationEnv.from_docker_image(IMAGE_NAME)216 217    all_results: List[Dict[str, Any]] = []218 219    try:220        for task in TASKS:221            result = await run_task(env=env, task_index=task["index"], task_name=task["name"])222            all_results.append(result)223            print(flush=True)224 225    finally:226        try:227            await env.close()228        except Exception as e:229            print(f"[DEBUG] env.close() error: {e}", file=sys.stderr, flush=True)230 231    total = sum(r["score"] for r in all_results) / len(all_results)232    print("=" * 60, flush=True)233    print(f"OVERALL SCORE: {total:.4f}", flush=True)234    for r in all_results:235        tag = "PASS" if r["success"] else "FAIL"236        print(f"  [{tag}] {r['task_name']:25s} score={r['score']:.4f}", flush=True)237 238 239if __name__ == "__main__":240    asyncio.run(main())241