negotiation-arena-master / inference.py
the-ashutosh's picture
Add application file
1a91b23
Raw History Blame Contribute Delete
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())