CoolFace
Apppublic

sc-likes-to-code/openenv-customer-support-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
inference.py144 linesDownload Raw Back to root
1import json2import os3from typing import List, Optional, Tuple4 5from openai import OpenAI6 7from server.your_environment import SupportEnv8from models import Action9 10# ── Environment / model config ──────────────────────────────────────────────11HF_TOKEN     = os.getenv("HF_TOKEN")12API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")13MODEL_NAME   = os.getenv("MODEL_NAME",   "Qwen/Qwen2.5-72B-Instruct")14BENCHMARK    = os.getenv("MY_ENV_V4_BENCHMARK", "support_env")15 16MAX_STEPS = 817TEMPERATURE = 0.018MAX_TOKENS = 15019 20ALL_TASKS = ["easy", "medium", "hard"]21 22# ── Logging helpers ──────────────────────────────────────────────────────────23def log_start(task: str, env: str, model: str) -> None:24    print(f"[START] task={task} env={env} model={model}", flush=True)25 26def log_step(step: int, action: str, reward: float, done: bool) -> None:27    print(f"[STEP] step={step} action={action} reward={reward:.2f} done={str(done).lower()} error=null", flush=True)28 29def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:30    rewards_str = ",".join(f"{r:.2f}" for r in rewards)31    print(f"[END] success={str(success).lower()} steps={steps} score={score:.2f} rewards={rewards_str}", flush=True)32 33# ── Fallback actions ─────────────────────────────────────────────────────────34def fallback_action(step, task, ticket_text, ticket_id):35    text = ticket_text.lower()36 37    is_billing = any(w in text for w in ["charged", "payment", "refund", "deducted"])38    39    if step == 1:40        return Action(41            action_type="classify",42            ticket_id=ticket_id,43            content="billing" if is_billing else "technical"44        )45 46    if step == 2:47        if task == "hard":48            return Action(49                action_type="ask",50                ticket_id=ticket_id,51                content="Please provide your transaction ID so I can check the payment."52            )53 54        if is_billing:55            return Action(56                action_type="respond",57                ticket_id=ticket_id,58                content="We are sorry for the issue. Your refund will be processed immediately."59            )60 61        return Action(62            action_type="respond",63            ticket_id=ticket_id,64            content="We are sorry for the inconvenience. We will investigate and fix the issue."65        )66 67    # step 3+68    if is_billing:69        return Action(70            action_type="respond",71            ticket_id=ticket_id,72            content="Thanks for the details. Your refund has been successfully processed."73        )74 75    return Action(76        action_type="respond",77        ticket_id=ticket_id,78        content="The issue has been fixed. Please check again."79    )80 81# ── Model action ─────────────────────────────────────────────────────────────82def get_model_action(client, step, task, ticket_text, ticket_id):83    if client is None:84        return fallback_action(step, task, ticket_text, ticket_id)85 86    try:87        client.chat.completions.create(88            model=MODEL_NAME,89            messages=[{"role": "user", "content": ticket_text}],90            temperature=TEMPERATURE,91            max_tokens=MAX_TOKENS,92        )93        return fallback_action(step, task, ticket_text, ticket_id)94    except:95        return fallback_action(step, task, ticket_text, ticket_id)96 97# ── Ticket extraction ────────────────────────────────────────────────────────98def extract_ticket(observation: dict):99    ticket = observation["tickets"][0]100    return ticket["id"], ticket["text"]101 102# ── Run single task ──────────────────────────────────────────────────────────103def run_task(client, task_name: str):104    env = SupportEnv()105    observation = env.reset(task_name).model_dump()106 107    rewards: List[float] = []108    steps_taken = 0109 110    log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)111 112    ticket_id, ticket_text = extract_ticket(observation)113 114    for step in range(1, MAX_STEPS + 1):115        action = get_model_action(client, step, task_name, ticket_text, ticket_id)116 117        observation, reward, done, _ = env.step(action)118        observation = observation.model_dump()119 120        ticket_id, ticket_text = extract_ticket(observation)121 122        reward_val = round(float(reward.score), 2)123        rewards.append(reward_val)124        steps_taken = step125 126        log_step(step=step, action=action.action_type, reward=reward_val, done=done)127 128        if done:129            break130 131    score = round(sum(rewards) / len(rewards), 2) if rewards else 0.00132    success = score >= 0.30133 134    log_end(success=success, steps=steps_taken, score=score, rewards=rewards)135 136# ── Main ─────────────────────────────────────────────────────────────────────137def main() -> None:138    client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN) if HF_TOKEN else None139 140    for task_name in ALL_TASKS:141        run_task(client, task_name)142 143if __name__ == "__main__":144    main()