CoolFace
Apppublic

Ar-Srivas/BitWise_CSS_env

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
inference.py500 linesDownload Raw Back to root
1"""2Baseline inference agent for css_env.3 4Uses the official OpenEnv EnvClient SDK (via the CssEnv client class5defined in client.py at project root) to drive the environment.6 7Reads environment variables:8  - API_BASE_URL  (default: https://api.openai.com/v1)9  - MODEL_NAME    (default: gpt-4o-mini)10  - HF_TOKEN or API_KEY or OPENAI_API_KEY (required)11 12Optional:13  - LOCAL_IMAGE_NAME  - Docker image name. If set, the env is launched via14                        CssEnv.from_docker_image(LOCAL_IMAGE_NAME).15  - ENV_URL           - direct base URL of an already-running css_env server16  - TASK_NAME         - task1/task2/task3/task4/easy/medium/hard/all17  - MAX_STEPS         - hard cap on loop steps per task18  - TEMPERATURE       - LLM sampling temperature19  - MAX_TOKENS        - LLM max output tokens20 21Emits the OpenEnv structured stdout format:22  [START] task=<task_name> env=<benchmark> model=<model_name>23  [STEP]  step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>24  [END]   success=<true|false> steps=<n> score=<0.000> rewards=<r1,r2,...,rn>25"""26 27from __future__ import annotations28 29import asyncio30import json31import os32import re33import sys34from typing import Any, Dict, List, Optional35 36from openai import OpenAI37 38from client import CssEnv39from models import CssAction40 41try:42    from server.tasks import TASKS, TASK_ORDER43except ImportError:44    from tasks import TASKS, TASK_ORDER45 46 47# Configuration48API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1").strip()49MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o-mini").strip()50LLM_API_KEY = (51    os.getenv("HF_TOKEN")52    or os.getenv("API_KEY")53    or os.getenv("OPENAI_API_KEY")54    or ""55).strip()56 57LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME", "").strip()58ENV_URL = os.getenv("ENV_URL", "http://localhost:8000").strip()59TASK_NAME = os.getenv("TASK_NAME", "all").strip()60MAX_STEPS = max(1, int(os.getenv("MAX_STEPS", "20")))61TEMPERATURE = float(os.getenv("TEMPERATURE", "0"))62MAX_TOKENS = max(64, int(os.getenv("MAX_TOKENS", "350")))63BENCHMARK = "css_env"64SCORE_EPSILON = 1e-665MIN_SCORE_BOUND = 0.0166MAX_SCORE_BOUND = 0.9967 68SYSTEM_PROMPT = """You are an expert frontend engineer fixing CSS to match design tokens.69 70Return ONLY one JSON object with this schema:71{72  "action_type": "replace_color|fix_spacing|fix_typography|fix_contrast|add_breakpoint|remove_rule",73  "target": "string",74  "value": "string or null"75}76 77Guidelines:781. Prioritize the lowest score in observation.scores.792. Do not repeat the same action signature.803. Prefer targeted edits that change CSS and improve one dimension.814. If a previous action had no effect, choose a different action type.825. No markdown, no explanation, JSON only.83"""84 85 86def clamp01(value: float) -> float:87    try:88        numeric = float(value)89    except (TypeError, ValueError):90        numeric = MIN_SCORE_BOUND + SCORE_EPSILON91 92    if numeric <= MIN_SCORE_BOUND:93        numeric = MIN_SCORE_BOUND + SCORE_EPSILON94    if numeric >= MAX_SCORE_BOUND:95        numeric = MAX_SCORE_BOUND - SCORE_EPSILON96 97    rounded = round(numeric, 2)98    if rounded <= MIN_SCORE_BOUND:99        return 0.02100    if rounded >= MAX_SCORE_BOUND:101        return 0.98102    return float(f"{rounded:.2f}")103 104 105def log_start(task: str, env: str, model: str) -> None:106    print(f"[START] task={task} env={env} model={model}", flush=True)107 108 109def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str] = None) -> None:110    done_str = "true" if done else "false"111    error_str = error if error else "null"112    print(113        f"[STEP] step={step} action={action} reward={reward:.2f} done={done_str} error={error_str}",114        flush=True,115    )116 117 118def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:119    success_str = "true" if success else "false"120    rewards_str = ",".join(f"{r:.2f}" for r in rewards)121    print(122        f"[END] success={success_str} steps={steps} score={clamp01(score):.3f} rewards={rewards_str}",123        flush=True,124    )125 126 127def _task_description(task_name: str) -> str:128    cfg = TASKS.get(task_name, {})129    return str(cfg.get("description") or cfg.get("difficulty") or task_name)130 131 132def _select_tasks(task_name: str) -> List[str]:133    aliases = {"easy": "task1", "medium": "task2", "hard": "task3"}134    requested = aliases.get(task_name.lower(), task_name.lower())135 136    if requested == "all":137        return [name for name in TASK_ORDER if name in TASKS]138    if requested in TASKS:139        return [requested]140 141    known = ",".join(["all", "easy", "medium", "hard", *TASKS.keys()])142    raise ValueError(f"Unknown TASK_NAME '{task_name}'. Expected one of: {known}")143 144 145def _extract_selectors(css: str) -> List[str]:146    selectors: List[str] = []147    for selector_group in re.findall(r"([^{}]+)\{", css or ""):148        for raw in selector_group.split(","):149            selector = raw.strip()150            if selector and selector not in selectors:151                selectors.append(selector)152    return selectors153 154 155def _extract_colors(css: str) -> List[str]:156    return re.findall(r"#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{6})", css or "")157 158 159def _first_selector(css: str) -> str:160    selectors = _extract_selectors(css)161    return selectors[0] if selectors else ".container"162 163 164def _build_user_prompt(observation: Any, step: int, task_name: str) -> str:165    html = str(getattr(observation, "html", ""))166    css = str(getattr(observation, "css", ""))167    tokens = getattr(observation, "tokens", {}) or {}168    violations = getattr(observation, "violations", None)169    scores = getattr(observation, "scores", {}) or {}170    prev_score = getattr(observation, "score", None)171 172    html_preview = html[:3000]173    css_preview = css[:5000]174 175    return (176        f"Task: {task_name} - {_task_description(task_name)}\n"177        f"Step: {step}\n"178        f"Current score: {prev_score}\n"179        f"Scores: {json.dumps(scores, separators=(',', ':'))}\n"180        f"Violations: {json.dumps(violations)}\n"181        f"Tokens: {json.dumps(tokens, separators=(',', ':'))}\n\n"182        f"HTML:\n{html_preview}\n\n"183        f"CSS:\n{css_preview}\n"184    )185 186 187def _parse_action_json(text: str) -> Optional[Dict[str, Any]]:188    raw = (text or "").strip()189    if not raw:190        return None191 192    if "```" in raw:193        lines = raw.splitlines()194        capture: List[str] = []195        in_block = False196        for line in lines:197            if line.strip().startswith("```"):198                if in_block:199                    break200                in_block = True201                continue202            if in_block:203                capture.append(line)204        if capture:205            raw = "\n".join(capture).strip()206 207    start = raw.find("{")208    end = raw.rfind("}")209    if start < 0 or end <= start:210        return None211 212    try:213        payload = json.loads(raw[start : end + 1])214    except json.JSONDecodeError:215        return None216 217    return payload if isinstance(payload, dict) else None218 219 220def _fallback_action(observation: Any) -> Dict[str, Any]:221    css = str(getattr(observation, "css", ""))222    tokens = getattr(observation, "tokens", {}) or {}223    selector = _first_selector(css)224 225    token_colors = list((tokens.get("colors") or {}).values())226    colors = _extract_colors(css)227    if colors and token_colors:228        replacement = str(token_colors[0])229        if replacement != colors[0]:230            return {231                "action_type": "replace_color",232                "target": colors[0],233                "value": replacement,234            }235 236    spacing = (tokens.get("spacing") or {}).get("md", 16)237    return {238        "action_type": "fix_spacing",239        "target": f"{selector}.margin",240        "value": f"{int(spacing)}px",241    }242 243 244def _normalize_action(action: Dict[str, Any], observation: Any) -> Dict[str, Any]:245    allowed = {246        "replace_color",247        "fix_spacing",248        "fix_typography",249        "fix_contrast",250        "add_breakpoint",251        "remove_rule",252    }253 254    safe = dict(action or {})255    action_type = str(safe.get("action_type", "")).strip()256    if action_type not in allowed:257        return _fallback_action(observation)258 259    css = str(getattr(observation, "css", ""))260    selectors = _extract_selectors(css)261    first_selector = selectors[0] if selectors else ".container"262 263    target = str(safe.get("target", "")).strip()264    value = safe.get("value", None)265 266    if action_type == "replace_color":267        colors = _extract_colors(css)268        if not target or target not in colors:269            target = colors[0] if colors else "#333333"270        if value is None or not str(value).strip():271            value = "#1a6fe0"272        return {"action_type": action_type, "target": target, "value": str(value)}273 274    if action_type in {"fix_spacing", "fix_typography"}:275        if "." not in target:276            prop = "margin" if action_type == "fix_spacing" else "font-size"277            target = f"{first_selector}.{prop}"278        if value is None or not str(value).strip():279            value = "16px"280        return {"action_type": action_type, "target": target, "value": str(value)}281 282    if action_type == "fix_contrast":283        if not target:284            target = first_selector285        if value is None or "," not in str(value):286            value = "#333333,#ffffff"287        return {"action_type": action_type, "target": target, "value": str(value)}288 289    if action_type == "add_breakpoint":290        if not re.fullmatch(r"\d+px", target):291            target = "768px"292        if value is None or "{" not in str(value):293            value = f"{first_selector} {{ width: 100%; }}"294        return {"action_type": action_type, "target": target, "value": str(value)}295 296    if action_type == "remove_rule":297        if not target:298            target = selectors[-1] if selectors else ".unused"299        return {"action_type": action_type, "target": target, "value": None}300 301    return _fallback_action(observation)302 303 304def _action_str(action: Dict[str, Any]) -> str:305    return json.dumps(action, separators=(",", ":"), ensure_ascii=True)306 307 308def _make_openai_client() -> OpenAI:309    if not LLM_API_KEY:310        raise ValueError("Missing API key. Set HF_TOKEN or API_KEY or OPENAI_API_KEY.")311    return OpenAI(base_url=API_BASE_URL, api_key=LLM_API_KEY)312 313 314def _probe_llm(client: OpenAI) -> None:315    client.chat.completions.create(316        model=MODEL_NAME,317        messages=[318            {"role": "system", "content": "Reply with exactly ok"},319            {"role": "user", "content": "ok"},320        ],321        temperature=0,322        max_tokens=2,323    )324 325 326async def _init_env() -> CssEnv:327    if LOCAL_IMAGE_NAME:328        return await CssEnv.from_docker_image(LOCAL_IMAGE_NAME)329 330    env = CssEnv(base_url=ENV_URL)331    await env.connect()332    return env333 334 335def _task_step_limit(task_cfg: Dict[str, Any]) -> int:336    task_limit = int(task_cfg.get("max_steps", MAX_STEPS) or MAX_STEPS)337    return max(1, min(task_limit, MAX_STEPS))338 339 340def _task_threshold(task_cfg: Dict[str, Any]) -> float:341    return clamp01(float(task_cfg.get("success_threshold", 0.95)))342 343 344def _llm_action(client: OpenAI, observation: Any, step: int, task_name: str) -> Dict[str, Any]:345    prompt = _build_user_prompt(observation, step, task_name)346    completion = client.chat.completions.create(347        model=MODEL_NAME,348        messages=[349            {"role": "system", "content": SYSTEM_PROMPT},350            {"role": "user", "content": prompt},351        ],352        temperature=TEMPERATURE,353        max_tokens=MAX_TOKENS,354    )355    text = completion.choices[0].message.content or ""356    parsed = _parse_action_json(text)357    if parsed is None:358        return _fallback_action(observation)359    return _normalize_action(parsed, observation)360 361 362async def run_task(env: CssEnv, llm_client: OpenAI, task_name: str, task_cfg: Dict[str, Any]) -> Dict[str, Any]:363    log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)364 365    rewards: List[float] = []366    final_score = 0.0367    steps_taken = 0368    final_success = False369 370    try:371        reset_result = await env.reset(task=task_cfg, seed=7)372        observation = reset_result.observation373        done = bool(reset_result.done)374    except Exception as exc:375        log_step(step=0, action="{}", reward=0.0, done=True, error=f"reset_failed:{exc}")376        log_end(success=False, steps=0, score=0.0, rewards=[])377        return {"score": 0.0, "steps": 0, "rewards": []}378 379    max_steps = _task_step_limit(task_cfg)380    threshold = _task_threshold(task_cfg)381 382    for step in range(1, max_steps + 1):383        if done:384            break385 386        llm_error: Optional[str] = None387        try:388            action_payload = _llm_action(llm_client, observation, step, task_name)389        except Exception as exc:390            llm_error = f"llm_error:{exc}"391            action_payload = _fallback_action(observation)392 393        action_payload = _normalize_action(action_payload, observation)394        action_text = _action_str(action_payload)395 396        step_error: Optional[str] = llm_error397        try:398            step_result = await env.step(CssAction(**action_payload))399        except Exception as exc:400            step_error = f"step_failed:{exc}"401            action_payload = _fallback_action(observation)402            action_text = _action_str(action_payload)403            try:404                step_result = await env.step(CssAction(**action_payload))405            except Exception as exc2:406                log_step(step=step, action=action_text, reward=0.0, done=True, error=f"step_failed:{exc2}")407                steps_taken = step408                break409 410        reward = clamp01(float(step_result.reward or 0.0))411        done = bool(step_result.done)412        observation = step_result.observation413        rewards.append(reward)414        steps_taken = step415 416        observed_score = getattr(observation, "score", None)417        if observed_score is not None:418            final_score = clamp01(float(observed_score))419        elif rewards:420            final_score = clamp01(sum(rewards) / len(rewards))421 422        final_success = bool(getattr(observation, "success", False)) or (final_score >= threshold)423        log_step(step=step, action=action_text, reward=reward, done=done, error=step_error)424 425    if not rewards:426        final_score = 0.0427 428    log_end(success=final_success, steps=steps_taken, score=final_score, rewards=rewards)429    return {"score": final_score, "steps": steps_taken, "rewards": rewards}430 431 432async def main_async() -> None:433    try:434        tasks_to_run = _select_tasks(TASK_NAME)435    except Exception as exc:436        print(f"[DEBUG] Invalid task selection: {exc}", file=sys.stderr, flush=True)437        return438 439    try:440        llm_client = _make_openai_client()441        _probe_llm(llm_client)442    except Exception as exc:443        for task_name in tasks_to_run:444            log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)445            log_step(step=0, action="{}", reward=0.0, done=True, error=f"llm_init_failed:{exc}")446            log_end(success=False, steps=0, score=0.0, rewards=[])447        return448 449    env: Optional[CssEnv] = None450    results: Dict[str, Dict[str, Any]] = {}451 452    try:453        env = await _init_env()454    except Exception as exc:455        for task_name in tasks_to_run:456            log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)457            log_step(step=0, action="{}", reward=0.0, done=True, error=f"env_init_failed:{exc}")458            log_end(success=False, steps=0, score=0.0, rewards=[])459        return460 461    try:462        for task_name in tasks_to_run:463            task_cfg = dict(TASKS.get(task_name, {}))464            if not task_cfg:465                log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)466                log_step(step=0, action="{}", reward=0.0, done=True, error="missing_task_config")467                log_end(success=False, steps=0, score=0.0, rewards=[])468                continue469 470            try:471                results[task_name] = await run_task(env, llm_client, task_name, task_cfg)472            except Exception as exc:473                log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)474                log_step(step=0, action="{}", reward=0.0, done=True, error=f"task_failed:{exc}")475                log_end(success=False, steps=0, score=0.0, rewards=[])476                results[task_name] = {"score": 0.0, "steps": 0, "rewards": []}477    finally:478        if env is not None:479            try:480                await env.close()481            except Exception as exc:482                print(f"[DEBUG] env.close() error: {exc}", file=sys.stderr, flush=True)483 484    print("\n=== Final Results ===", file=sys.stderr)485    total = 0.0486    for task_name in tasks_to_run:487        score = clamp01(float(results.get(task_name, {}).get("score", 0.0)))488        total += score489        print(f"  {task_name}: score={score:.4f}", file=sys.stderr)490    if tasks_to_run:491        print(f"  average: {total / len(tasks_to_run):.4f}", file=sys.stderr)492 493 494def main() -> None:495    asyncio.run(main_async())496 497 498if __name__ == "__main__":499    main()500