Download inference.py from ashutosh111/negotiation-arena-master: direct link, hf CLI and curl.
- Browser
- Download file 8.04 kB
-
https://huggingface.co/spaces/ashutosh111/negotiation-arena-master/resolve/main/inference.py
- Command line
-
hf download hf://spaces/ashutosh111/negotiation-arena-master/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/ashutosh111/negotiation-arena-master/resolve/main/inference.py
8.04 kB
| """Negotiation Arena — Vendor inference driver (V3). | |
| Runs all 3 tasks against the (running) HTTP server using a Llama-3.3-70B | |
| Vendor agent. The Vendor produces a Theory-of-Mind prediction alongside | |
| each action; the server grades the prediction and may run a 3-judge | |
| Tribunal at episode end. This script reads those signals back and prints | |
| them as optional fields on the existing log lines. | |
| Output format (V2-compatible; new fields are optional): | |
| [START] task=<task_id> env=negotiation-arena model=<MODEL_NAME> | |
| [STEP] step=<n> action=<json>[ tom_acc=<0.00>] reward=<0.00> done=<true|false> error=<msg|null> | |
| [END] success=<true|false> steps=<n> score=<0.000>[ tribunal=<0.000>] rewards=<r1,r2,...,rn> | |
| The ``tom_acc=`` field appears only when the server returned a graded ToM | |
| prediction for that turn. The ``tribunal=`` field appears only when the | |
| server's terminal observation included a tribunal breakdown. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import os | |
| import sys | |
| from typing import Any, Dict, List, Optional | |
| from client import NegotiationArenaEnv | |
| from evaluation.llm_vendor import LLMVendor | |
| # --------------------------------------------------------------------------- | |
| # Environment knobs (V2-compatible — same names, same defaults) | |
| # --------------------------------------------------------------------------- | |
| API_BASE_URL = os.environ.get("API_BASE_URL", "https://router.huggingface.co/v1") | |
| MODEL_NAME = os.environ.get("MODEL_NAME", "meta-llama/Llama-3.3-70B-Instruct") | |
| HF_TOKEN = os.environ.get("HF_TOKEN", "") | |
| SPACE_URL = os.environ.get("SPACE_URL", "http://localhost:8000") | |
| EVAL_MODE = os.environ.get("EVAL_MODE", "false").lower() == "true" | |
| TASKS = ["simple_saas", "gdpr_dpa", "enterprise_partnership"] | |
| # --------------------------------------------------------------------------- | |
| # Logging helpers | |
| # --------------------------------------------------------------------------- | |
| # log_start / log_step / log_end keep their V2 signatures and produce identical | |
| # output when no new V3 fields are passed. New optional kwargs add the fields. | |
| def log_start(task_id: str) -> None: | |
| print( | |
| f"[START] task={task_id} env=negotiation-arena model={MODEL_NAME}", | |
| flush=True, | |
| ) | |
| def log_step( | |
| step: int, | |
| action: Dict[str, Any], | |
| reward: float, | |
| done: bool, | |
| error: Optional[str], | |
| tom_acc: Optional[float] = None, | |
| ) -> None: | |
| """Print a [STEP] line. | |
| When ``tom_acc`` is provided, an additional ``tom_acc=X.XX`` field is | |
| inserted between ``action=...`` and ``reward=...``. Otherwise the line | |
| is byte-for-byte identical to V2. | |
| """ | |
| err = "null" if not error else error | |
| action_json = json.dumps(_compact_action(action), separators=(",", ":")) | |
| tom_field = f" tom_acc={float(tom_acc):.2f}" if tom_acc is not None else "" | |
| print( | |
| f"[STEP] step={step} action={action_json}{tom_field} " | |
| f"reward={reward:.2f} done={'true' if done else 'false'} error={err}", | |
| flush=True, | |
| ) | |
| def log_end( | |
| success: bool, | |
| steps: int, | |
| score: float, | |
| rewards: List[float], | |
| tribunal_score: Optional[float] = None, | |
| ) -> None: | |
| """Print an [END] line. | |
| When ``tribunal_score`` is provided, a ``tribunal=X.XXX`` field is inserted | |
| between ``score=...`` and ``rewards=...``. | |
| """ | |
| rewards_str = ",".join(f"{r:.2f}" for r in rewards) | |
| tribunal_field = ( | |
| f" tribunal={float(tribunal_score):.3f}" if tribunal_score is not None else "" | |
| ) | |
| print( | |
| f"[END] success={'true' if success else 'false'} steps={steps} " | |
| f"score={score:.3f}{tribunal_field} rewards={rewards_str}", | |
| flush=True, | |
| ) | |
| def _compact_action(action: Dict[str, Any]) -> Dict[str, Any]: | |
| """Trim large fields from the logged action.""" | |
| out: Dict[str, Any] = {} | |
| for k in ("action_type", "issue_name", "new_value"): | |
| if action.get(k) is not None: | |
| out[k] = action[k] | |
| if action.get("proposed_terms"): | |
| out["proposed_terms"] = action["proposed_terms"] | |
| if action.get("reasoning"): | |
| out["reasoning"] = str(action["reasoning"])[:120] | |
| return out | |
| # --------------------------------------------------------------------------- | |
| # Per-task driver | |
| # --------------------------------------------------------------------------- | |
| def run_task(env: NegotiationArenaEnv, vendor: LLMVendor, task_id: str) -> bool: | |
| log_start(task_id) | |
| rewards: List[float] = [] | |
| tribunal_score: Optional[float] = None | |
| success = False | |
| step_count = 0 | |
| final_score = 0.0 | |
| obs = env.reset(task_id=task_id, eval_mode=EVAL_MODE, client_strategy="balanced") | |
| max_steps = int(obs.get("max_turns", 30)) + 2 # safety margin | |
| while True: | |
| if obs.get("done"): | |
| break | |
| if step_count >= max_steps: | |
| break | |
| # V3: Vendor produces (tom_prediction, action) in one LLM call. | |
| tom_pred, action = vendor.act_with_tom(obs, turn=step_count + 1) | |
| # Stamp the boilerplate fields the server expects. | |
| action.setdefault("episode_id", obs.get("episode_id", "")) | |
| action.setdefault("turn_number", int(obs.get("turn_number", 0)) + 1) | |
| action.setdefault("agent_role", "vendor") | |
| # Attach the ToM prediction (server treats this as optional — V2-compatible). | |
| if tom_pred: | |
| action["tom_prediction"] = tom_pred | |
| try: | |
| obs = env.step(action) | |
| error = None | |
| except Exception as e: # noqa: BLE001 | |
| error = str(e) | |
| log_step(step_count + 1, action, 0.0, True, error) | |
| break | |
| step_count += 1 | |
| reward = float(obs.get("reward") or 0.0) | |
| rewards.append(reward) | |
| # ToM accuracy is in the observation when the server graded a prediction | |
| # this turn. Field is None on V2 servers or when no prediction was sent. | |
| tom_acc: Optional[float] = None | |
| feedback = obs.get("tom_grader_feedback") | |
| if isinstance(feedback, dict) and "composite" in feedback: | |
| try: | |
| tom_acc = float(feedback["composite"]) | |
| except (TypeError, ValueError): | |
| tom_acc = None | |
| log_step( | |
| step_count, | |
| action, | |
| reward, | |
| bool(obs.get("done")), | |
| obs.get("error_message"), | |
| tom_acc=tom_acc, | |
| ) | |
| if obs.get("done"): | |
| # Tribunal breakdown is only present on the terminal observation, | |
| # and only when the server's tribunal feature is enabled and an | |
| # API key is configured server-side. | |
| breakdown = obs.get("tribunal_breakdown") | |
| if isinstance(breakdown, dict) and "trimmed_mean" in breakdown: | |
| try: | |
| tribunal_score = float(breakdown["trimmed_mean"]) | |
| except (TypeError, ValueError): | |
| tribunal_score = None | |
| break | |
| if rewards: | |
| # Score = average per-turn reward (final entry is the composite if terminal). | |
| final_score = sum(rewards) / len(rewards) | |
| success = bool(obs.get("done")) and bool(obs.get("final_deal_struck", False)) | |
| log_end(success, step_count, final_score, rewards, tribunal_score=tribunal_score) | |
| return success | |
| def main() -> int: | |
| vendor = LLMVendor( | |
| api_base_url=API_BASE_URL, | |
| model_name=MODEL_NAME, | |
| api_key=HF_TOKEN, | |
| temperature=0.2, | |
| max_tokens=800, | |
| ) | |
| env = NegotiationArenaEnv(base_url=SPACE_URL) | |
| try: | |
| try: | |
| env.health() | |
| except Exception as e: # noqa: BLE001 | |
| print(f"[FATAL] cannot reach server at {SPACE_URL}: {e}", flush=True) | |
| return 1 | |
| for task_id in TASKS: | |
| try: | |
| run_task(env, vendor, task_id) | |
| except Exception as e: # noqa: BLE001 | |
| print(f"[ERROR] task {task_id} crashed: {e}", flush=True) | |
| finally: | |
| env.close() | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |