CoolFace
Apppublic

StarTripper/ticket_ordering

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py206 linesDownload Raw Back to root
1import os2import asyncio3import textwrap4import numpy as np5from openai import OpenAI6from typing import List, Optional, Dict, Any7 8from client import TicketOrderingEnv9from problem_generator import GenerationDifficulty10from models import TicketOrderingAction, TicketOrderingObservation11 12 13backup_rng = np.random.default_rng(42)14 15 16ENV_BASE = "https://startripper-openenv-ticket-ordering-env.hf.space"17 18API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY")19API_BASE_URL = os.getenv("API_BASE_URL", "https://api.groq.com/openai/v1")20MODEL_NAME = os.getenv("MODEL_NAME") or "llama-3.1-8b-instant"21MAX_STEPS = 1522TEMPERATURE = 0.0 # Kind of makes these models TOO deterministic / repeat things, oh well, rules are rules.23MAX_TOKENS = 30024SUCCESS_SCORE_THRESHOLD = 0.7525 26SYSTEM_PROMPT = textwrap.dedent(27    """28    You are a ticket prioritization agent.29 30    Given:31    - ordering criteria32    - reference tickets33    - a candidate ticket34    - heuristics from n of the most assigned and n of the least assigned tickets35    - total number of tickets36    - iterations completed so far37 38    Your job:39    - Assign a priority score to the candidate (relative to the references, must NOT be neutral / 0.0) (float, higher = more important)40    - Write a short summary for the candidate (<=32 chars) (Part of that ticket's heuristic)41    - Select next reference ticket ids (must be one of the keys from the heuristics)42    - Select next candidate ticket id (must be one of the keys from the heuristics)43    - Decide whether to end ordering (there is no need to end after iterations completed = total tickets since cross comparing tickets still may be valuable)44 45    Respond with no fancy formatting, backticks or anything else, respond ONLY with PURE valid JSON in this format:46    {47      "candidate_priority": float,48      "candidate_summary": string,49      "next_reference_ids": [int],50      "next_candidate_id": int,51      "end_ordering": bool52    }53    """54).strip()55 56 57def log_start(task: str, env: str, model: str) -> None:58    print(f"[START] task={task} env={env} model={model}", flush=True)59 60 61def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:62    error_val = error if error else "null"63    done_val = str(done).lower()64    print(65        f"[STEP] step={step} action={action} reward={reward:.2f} done={done_val} error={error_val}",66        flush=True,67    )68 69 70def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:71    rewards_str = ",".join(f"{r:.2f}" for r in rewards)72    print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)73 74 75def serialize_ticket(ticket: Any) -> Dict[str, Any]:76    return {77        "id": ticket.id,78        "thread": [{"user": c.user, "content": c.content} for c in ticket.thread],79        "heuristic": {80            "priority": ticket.heuristic.priority,81            "summary": ticket.heuristic.summary,82            "times_assigned": ticket.heuristic.times_assigned,83        },84    }85 86 87def build_user_prompt(obs: Any) -> str:88    return textwrap.dedent(89        f"""90        Ordering criteria: {obs.ordering_criteria}91 92        Reference tickets:93        {[serialize_ticket(t) for t in obs.reference_tickets]}94 95        Candidate ticket:96        {serialize_ticket(obs.candidate_ticket)}97 98        Existing heuristics:99        { {k: v.model_dump() for k, v in obs.ticket_heuristics.items()} }100 101        Total tickets: {obs.total_tickets}102        Completed iterations: {obs.completed_iterations}103 104        Decide the next action.105        """106    ).strip()107 108 109def get_model_action(client: OpenAI, obs: TicketOrderingObservation) -> Dict[str, Any]:110    user_prompt = build_user_prompt(obs)111 112    try:113        completion = client.chat.completions.create(114            model=MODEL_NAME,115            messages=[116                {"role": "system", "content": SYSTEM_PROMPT},117                {"role": "user", "content": user_prompt},118            ],119            temperature=TEMPERATURE,120            max_tokens=MAX_TOKENS,121        )122        text = (completion.choices[0].message.content or "").strip()123 124        import json125 126        return json.loads(text)127    except Exception:128        return {129            "candidate_priority": backup_rng.uniform(low=0.0, high=1.0),130            "candidate_summary": "issue",131            "next_reference_ids": [],132            "next_candidate_id": backup_rng.choice(list(obs.ticket_heuristics.keys())),133            "end_ordering": False,134        }135 136 137def main() -> None:138    client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)139    with TicketOrderingEnv(base_url=ENV_BASE).sync() as env:140        # async with TicketOrderingEnv(base_url="localhost:8001") as env:141        for task in [GenerationDifficulty.Easy, GenerationDifficulty.Medium, GenerationDifficulty.Hard]:142            rewards: List[float] = []143            steps_taken = 0144            score = 0.0145            success = False146 147            log_start(task=f"ticket_ordering_env_{task.name}", env="ticket_ordering_env", model=MODEL_NAME)148 149            try:150                result = env.reset(difficulty=task.value)151                obs = result.observation152 153                for step in range(1, MAX_STEPS + 1):154                    if result.done:155                        break156 157                    action_dict = get_model_action(client, obs)158 159                    action = TicketOrderingAction(160                        candidate_priority=float(action_dict.get("candidate_priority", 0.0)),161                        candidate_summary=str(action_dict.get("candidate_summary", ""))[:32],162                        next_reference_ids=list(action_dict.get("next_reference_ids", [])),163                        next_candidate_id=int(action_dict.get("next_candidate_id", 0)),164                        end_ordering=bool(action_dict.get("end_ordering", False)),165                    )166 167                    result = env.step(action)168                    obs = result.observation169 170                    reward = result.reward or 0.0171                    done = result.done172                    error = None173 174                    rewards.append(reward)175                    steps_taken = step176 177                    log_step(178                        step=step,179                        action=str(action_dict),180                        reward=reward,181                        done=done,182                        error=error,183                    )184 185                    if done:186                        break187 188                min_reward = -1.0189                max_reward = 2.0190                rewards_sum = sum(rewards)191                rewards_sum -= min_reward192                rewards_sum /= (max_reward - min_reward)193                score = min(max(rewards_sum, 0.0), 1.0)194                success = score >= SUCCESS_SCORE_THRESHOLD195 196            finally:197                try:198                    env.close()199                except Exception:200                    pass201                log_end(success=success, steps=steps_taken, score=score, rewards=rewards)202 203 204if __name__ == "__main__":205    main()206