Team Ai
Apppublic

Ar-Srivas/BitWise_CSS_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
css_env_environment.py623 linesDownload Raw Back to server
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7# make sure the doc strings are correct, i have done some check the others8 9"""10Css Env Environment Implementation.11"""12 13import re14import sys15import math16from pathlib import Path17from typing import Dict, List, Optional18from uuid import uuid419 20from openenv.core.env_server.interfaces import Environment21from openenv.core.env_server.types import State22 23_ROOT_DIR = Path(__file__).resolve().parents[1]24if str(_ROOT_DIR) not in sys.path:25    sys.path.insert(0, str(_ROOT_DIR))26 27try:28    from .models import CssAction, CssObservation, GradeResult29except ImportError:30    from models import CssAction, CssObservation, GradeResult31 32from reward import compute_reward33from graders import colors, spacing, typography, contrast, layout, cleanliness, design_quality34 35try:36    from .tasks import TASKS37    from .action_engine import apply_action38    from .flaw_injector import inject_flaws39except ImportError:40    from server.tasks import TASKS41    from server.action_engine import apply_action42    from server.flaw_injector import inject_flaws43 44 45class CssEnvironment(Environment):46    """47    CSS Refinement RL Environment48 49    This environment simulates a CSS refinement task where50    an agent iteratively modifies CSS code to improve design quality.51 52    Example:53    >>> env = CssEnvironment()54    >>> obs = env.reset(task, seed=42)55    >>> print(obs.html) # <div class="card">Hello</div>56    >>> print(obs.css) # .card { color: #1a73e7; }57    >>> result = env.step(CssAction(58    ...     action_type="replace_color",59    ...     target="#1a73e7",60    ...     value="#1a6fe0"61    ... ))62    >>> print(result["observation"].css) # .card { color: #1a6fe0; }63    >>> print(result["reward"]) # 0.664    >>> print(result["done"]) # False65    """66 67    # Enable concurrent WebSocket sessions.68    # Set to True if your environment isolates state between instances.69    # When True, multiple WebSocket clients can connect simultaneously, each70    # getting their own environment instance (when using factory mode in app.py).71    SUPPORTS_CONCURRENT_SESSIONS: bool = True72    HIGH_SCORE_GUARD_THRESHOLD: float = 0.9073    HIGH_SCORE_MAX_DROP: float = 0.0374    SCORE_EPSILON: float = 1e-675    MIN_SCORE_BOUND: float = 0.0176    MAX_SCORE_BOUND: float = 0.9977    SUPPORTED_GRADERS = (78        "color",79        "spacing",80        "typography",81        "contrast",82        "layout",83        "cleanliness",84        "design_quality",85    )86 87    def __init__(self):88        """Initialize the css_env environment."""89        self.html = "" # html string in the task90        self.css = "" # css string in the task91        self.tokens = {} # design tokens given in the task as a dict of token_name -> value92        self.config = {} # config given in the task to inject flaws93        self.manifest = [] # list of dicts describing flaws94        self.difficulty = "easy"95        self.success_threshold = 0.9596        self.required_graders = ["color", "spacing", "typography", "contrast", "cleanliness"]97        self.grader_weights = {98            "color": 0.30,99            "spacing": 0.20,100            "typography": 0.20,101            "contrast": 0.20,102            "cleanliness": 0.10,103        }104        self.step_count = 0 # current step count105        self.max_steps = 20 # maximum steps106        self.state_data = {} # tracks progress107 108    @staticmethod109    def _action_signature(action: CssAction) -> str:110        value = "null" if action.value is None else str(action.value)111        return f"{action.action_type}|{action.target}|{value}"112 113    @staticmethod114    def _action_score_key(action_type: str) -> Optional[str]:115        mapping = {116            "replace_color": "color",117            "fix_spacing": "spacing",118            "fix_typography": "typography",119            "fix_contrast": "contrast",120            "add_breakpoint": "layout",121            "remove_rule": "cleanliness",122        }123        return mapping.get(action_type)124 125    def _is_irrelevant_action(self, action: CssAction, scores: Dict[str, float]) -> bool:126        # In easy tasks, breakpoint tuning is usually not part of the objective.127        if self.difficulty == "easy" and action.action_type == "add_breakpoint":128            return True129 130        score_key = self._action_score_key(action.action_type)131        if score_key and scores.get(score_key, 0.0) >= self.success_threshold:132            return True133 134        return False135 136    def reset(self, task: Optional[Dict] = None, seed: int = 0) -> CssObservation:137            """138            task = {139                "html": str,140                "css": str,141                "design_tokens": dict,142                "config": dict,143                "difficulty": "easy" | "medium" | "hard"144            }145            """146 147            self.step_count = 0148 149            if task is None or not isinstance(task, dict) or not task:150                task = TASKS["task1"]151 152            self.html = task["html"]153            self.tokens = task.get("design_tokens", task.get("tokens", {}))154            self.config = task.get("config", task.get("flaw_config", {}))155            self.difficulty = task.get("difficulty", "easy")156            self.max_steps = int(task.get("max_steps", self.max_steps))157            self.success_threshold = float(task.get("success_threshold", 0.95))158            task_weights = self._normalize_grader_weights(task.get("grader_weights", {}))159            task_graders = task.get("required_graders", task.get("graders", self.required_graders))160            required_graders = [161                str(key).strip() for key in (task_graders if isinstance(task_graders, list) else []) if str(key).strip()162            ]163            required_graders = [164                key for key in required_graders if key in self.SUPPORTED_GRADERS165            ]166            if not required_graders:167                required_graders = [168                    key for key, weight in task_weights.items() if weight > 0.0 and key in self.SUPPORTED_GRADERS169                ]170 171            if required_graders:172                self.required_graders = required_graders173            else:174                self.required_graders = ["color", "spacing", "typography", "contrast", "cleanliness"]175 176            self.grader_weights = task_weights or self._default_weights_for_graders(self.required_graders)177 178            clean_css = task.get("css", task.get("clean_css", ""))179 180            result = inject_flaws(clean_css, self.tokens, self.config, seed)181 182            self.css = result["css"]183            self.manifest = result["manifest"]184            initial_unused_selectors = self._compute_unused_selectors(self.css, self.html)185 186            if task.get("initial_unused_selectors"):187                initial_unused_selectors = sorted(188                    set(initial_unused_selectors) | set(task["initial_unused_selectors"])189                )190 191            self.state_data = {192                "episode_id": str(uuid4()),193                "step_count": 0,194                "initial_manifest": self.manifest,195                "initial_unused_selectors": initial_unused_selectors,196                "last_action_signature": None,197                "same_action_count": 0,198                "reward_drop_streak": 0,199                "last_scores": {},200                "last_reward": None,201                "last_done": False,202                "last_success": False,203                "last_changed": True,204                "last_no_op": False,205                "last_irrelevant": False,206                "action_counts": {},207                "required_graders": list(self.required_graders),208                "grader_weights": dict(self.grader_weights),209            }210 211            initial_grade = self.grade()212            self.state_data["last_scores"] = dict(initial_grade.breakdown)213            self.state_data["last_grade"] = float(initial_grade.score)214            self.state_data["peak_score"] = (215                min(self.state_data["last_scores"].values())216                if self.state_data.get("last_scores")217                else 0.0218            )219 220            violations = None221            if self.difficulty == "easy":222                violations = task.get("violations")223                if violations is None:224                    violations = [225                        {"type": "hint", "selector": "", "message": hint}226                        for hint in result.get("hints", [])227                    ]228 229            return CssObservation(230                html=self.html,231                css=self.css,232                tokens=self.tokens,233                violations=violations,234                scores=self.state_data.get("last_scores", {}),235                score=float(self.state_data.get("last_grade", 0.0)),236                success=False,237                changed=True,238                no_op_action=False,239                repeated_action=False,240                terminated_by_max_steps=False,241            )242 243    def step(self, action: CssAction) -> CssObservation:244        """245        Execute a step in the environment by applying the action generated by the agent to the CSS.246        """247        prev_scores = dict(self.state_data.get("last_scores", {}))248        prev_reward = self.state_data.get("last_reward")249        prev_score = float(self.state_data.get("last_grade", min(prev_scores.values()) if prev_scores else 0.0))250        prev_peak_score = float(self.state_data.get("peak_score", prev_score))251        old_css = self.css252        new_css, _ = apply_action(self.css, action)253        css_changed = old_css.strip() != new_css.strip()254        action_sig = self._action_signature(action)255        repeated_action = action_sig == self.state_data.get("last_action_signature")256        same_action_count = int(self.state_data.get("same_action_count", 0)) + 1 if repeated_action else 1257        action_counts = dict(self.state_data.get("action_counts", {}))258        action_counts[action_sig] = int(action_counts.get(action_sig, 0)) + 1259        duplicate_action = action_counts[action_sig] > 1260        no_op_action = not css_changed261 262        self.css = new_css263        grade_result = self.grade()264        scores = dict(grade_result.breakdown)265        success = self._is_done(scores)266        irrelevant_action = self._is_irrelevant_action(action, prev_scores)267 268        # Compute reward with progress and anti-repeat penalties.269        reward = compute_reward(270            scores,271            action_valid=css_changed,272            done=success,273            action_repeated=repeated_action,274            action_duplicate=duplicate_action,275            action_irrelevant=irrelevant_action,276            previous_scores=prev_scores,277            previous_reward=prev_reward,278        )279 280        # FIX 1: block repeated identical actions with explicit penalty.281        if repeated_action:282            reward = max(0.0, reward - 0.05)283 284        # Success must be driven by grader thresholds only.285        # End the episode immediately once all required scores meet the threshold.286        if success:287            self.step_count += 1288            score = float(grade_result.score)289 290            info = {291                "scores": scores,292                "step_count": self.step_count,293                "changed": css_changed,294                "no_op_action": no_op_action,295                "repeated_action": repeated_action,296                "same_action_count": same_action_count,297                "duplicate_action": duplicate_action,298                "irrelevant_action": irrelevant_action,299                "reward_drop_too_much": False,300                "reward_drop_streak": 0,301                "reward_drop_terminated": False,302                "repeated_action_terminated": False,303                "same_action_cap_reached": False,304                "success": True,305                "high_score_degradation": False,306                "terminated_by_max_steps": False,307                "success_threshold": self.success_threshold,308                "score": score,309            }310 311            self.state_data["step_count"] = self.step_count312            self.state_data["last_scores"] = scores313            self.state_data["last_grade"] = score314            self.state_data["last_reward"] = reward315            self.state_data["last_done"] = True316            self.state_data["last_success"] = True317            self.state_data["last_changed"] = css_changed318            self.state_data["last_no_op"] = no_op_action319            self.state_data["last_irrelevant"] = irrelevant_action320            self.state_data["last_action_signature"] = action_sig321            self.state_data["same_action_count"] = same_action_count322            self.state_data["reward_drop_streak"] = 0323            self.state_data["peak_score"] = max(prev_peak_score, score)324            self.state_data["action_counts"] = action_counts325 326            return CssObservation(327                html=self.html,328                css=self.css,329                tokens=self.tokens,330                violations=None,331                scores=scores,332                score=score,333                success=True,334                changed=css_changed,335                no_op_action=no_op_action,336                repeated_action=repeated_action,337                terminated_by_max_steps=False,338                reward=reward,339                done=True,340                metadata=info,341            )342 343        # Track large reward drops, but allow recovery attempts.344        reward_drop_too_much = (345            prev_reward is not None and reward < (float(prev_reward) - 0.2)346        )347        reward_drop_streak = (348            int(self.state_data.get("reward_drop_streak", 0)) + 1349            if reward_drop_too_much350            else 0351        )352        reward_drop_terminated = reward_drop_streak > 2353 354        # FIX 4: cap repeated identical action attempts.355        same_action_cap_reached = same_action_count > 2356        repeated_action_terminated = same_action_cap_reached357 358        self.step_count += 1359        done = (360            success361            or self.step_count >= self.max_steps362            or repeated_action_terminated363            or reward_drop_terminated364            or same_action_cap_reached365        )366        terminated_by_max_steps = done and not success367        score = float(grade_result.score)368 369        high_score_degradation = (370            prev_peak_score >= self.HIGH_SCORE_GUARD_THRESHOLD371            and score < (prev_peak_score - self.HIGH_SCORE_MAX_DROP)372        )373 374        if high_score_degradation:375            # Keep policy from drifting after reaching a strong state.376            self.css = old_css377            scores = prev_scores378            score = prev_score379            success = self._is_done(scores)380            done = True381            reward = max(0.0, (float(prev_reward) - 0.05) if prev_reward is not None else reward)382 383        info = {384            "scores": scores,385            "step_count": self.step_count,386            "changed": css_changed,387            "no_op_action": no_op_action,388            "repeated_action": repeated_action,389            "same_action_count": same_action_count,390            "duplicate_action": duplicate_action,391            "irrelevant_action": irrelevant_action,392            "reward_drop_too_much": reward_drop_too_much,393            "reward_drop_streak": reward_drop_streak,394            "reward_drop_terminated": reward_drop_terminated,395            "repeated_action_terminated": repeated_action_terminated,396            "same_action_cap_reached": same_action_cap_reached,397            "success": success,398            "high_score_degradation": high_score_degradation,399            "terminated_by_max_steps": terminated_by_max_steps,400            "success_threshold": self.success_threshold,401            "score": score,402        }403 404        self.state_data["step_count"] = self.step_count405        self.state_data["last_scores"] = scores406        self.state_data["last_grade"] = score407        self.state_data["last_reward"] = reward408        self.state_data["last_done"] = done409        self.state_data["last_success"] = success410        self.state_data["last_changed"] = css_changed411        self.state_data["last_no_op"] = no_op_action412        self.state_data["last_irrelevant"] = irrelevant_action413        self.state_data["last_action_signature"] = action_sig414        self.state_data["same_action_count"] = same_action_count415        self.state_data["reward_drop_streak"] = reward_drop_streak416        self.state_data["peak_score"] = max(prev_peak_score, score)417        self.state_data["action_counts"] = action_counts418 419        return CssObservation(420            html=self.html,421            css=self.css,422            tokens=self.tokens,423            violations=None,424            scores=scores,425            score=score,426            success=success,427            changed=css_changed,428            no_op_action=no_op_action,429            repeated_action=repeated_action,430            terminated_by_max_steps=terminated_by_max_steps,431            reward=reward,432            done=done,433            metadata=info,434        )435 436    @property437    def state(self) -> State:438        """439        Get the current environment state.440 441        Returns:442            Current State with episode_id and step_count443        """444        return State(445            episode_id=self.state_data.get("episode_id"),446            step_count=self.state_data.get("step_count", 0),447            success=self.state_data.get("last_success", False),448            done=self.state_data.get("last_done", False),449            reward=self.state_data.get("last_reward") if self.state_data.get("last_reward") is not None else 0.0,450            changed=self.state_data.get("last_changed", True),451            no_op_action=self.state_data.get("last_no_op", False),452            scores=self.state_data.get("last_scores", {}),453        )454 455    @staticmethod456    def _normalize_grader_weights(raw_weights: Dict[str, float]) -> Dict[str, float]:457        if not isinstance(raw_weights, dict):458            return {}459 460        normalized: Dict[str, float] = {}461        for grader_key, weight in raw_weights.items():462            key = str(grader_key).strip()463            if not key:464                continue465            try:466                normalized[key] = float(weight)467            except (TypeError, ValueError):468                continue469 470        return normalized471 472    @staticmethod473    def _default_weights_for_graders(graders: List[str]) -> Dict[str, float]:474        normalized = [str(key).strip() for key in graders if str(key).strip()]475        if not normalized:476            return {}477        equal_weight = 1.0 / float(len(normalized))478        return {key: equal_weight for key in normalized}479 480    def _active_grader_weights(self) -> Dict[str, float]:481        candidate = self.state_data.get("grader_weights", self.grader_weights)482        weights = {483            key: value484            for key, value in self._normalize_grader_weights(candidate).items()485            if key in self.SUPPORTED_GRADERS486        }487        if weights:488            return weights489        return self._default_weights_for_graders(list(self.required_graders))490 491    @classmethod492    def _clamp_open_unit_interval(cls, value: float) -> float:493        try:494            numeric = float(value)495        except (TypeError, ValueError):496            numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON497        if not math.isfinite(numeric):498            numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON499        if numeric <= cls.MIN_SCORE_BOUND:500            numeric = cls.MIN_SCORE_BOUND + cls.SCORE_EPSILON501        if numeric >= cls.MAX_SCORE_BOUND:502            numeric = cls.MAX_SCORE_BOUND - cls.SCORE_EPSILON503 504        rounded = round(numeric, 2)505        if rounded <= cls.MIN_SCORE_BOUND:506            return 0.02507        if rounded >= cls.MAX_SCORE_BOUND:508            return 0.98509        return float(f"{rounded:.2f}")510 511    def grade(self) -> GradeResult:512        """Grade the current CSS using task-configured grader weights."""513        all_scores = self._run_graders()514        weights = self._active_grader_weights()515 516        breakdown: Dict[str, float] = {}517        details: Dict[str, str] = {}518        weighted_sum = 0.0519        weight_sum = 0.0520 521        for grader_key, weight in weights.items():522            if weight <= 0.0:523                continue524 525            component_score = float(all_scores.get(grader_key, 0.0))526            breakdown[grader_key] = component_score527            weighted_sum += component_score * weight528            weight_sum += weight529            details[grader_key] = (530                f"component={grader_key}, score={component_score:.4f}, weight={weight:.4f}"531            )532 533        if not breakdown:534            required = list(self.required_graders) or ["color", "spacing", "typography", "contrast", "cleanliness"]535            for grader_key in required:536                component_score = float(all_scores.get(grader_key, 0.0))537                breakdown[grader_key] = component_score538                details[grader_key] = (539                    f"component={grader_key}, score={component_score:.4f}, weight=auto"540                )541            score = min(breakdown.values()) if breakdown else 0.0542            return GradeResult(score=self._clamp_open_unit_interval(score), breakdown=breakdown, details=details)543 544        score = (weighted_sum / weight_sum) if weight_sum > 0.0 else 0.0545        return GradeResult(score=self._clamp_open_unit_interval(score), breakdown=breakdown, details=details)546 547    def _run_graders(self) -> Dict[str, float]:548        return {549            "color": self._grade_color(),550            "spacing": self._grade_spacing(),551            "typography": self._grade_typography(),552            "contrast": self._grade_contrast(),553            "layout": self._grade_layout(),554            "cleanliness": self._grade_cleanliness(),555            "design_quality": self._grade_design_quality(),556        }557 558    def _grade_color(self) -> float:559        return self._safe_grade(colors.grade)560 561    def _grade_spacing(self) -> float:562        return self._safe_grade(spacing.grade)563 564    def _grade_typography(self) -> float:565        return self._safe_grade(typography.grade)566 567    def _grade_contrast(self) -> float:568        return self._safe_grade(contrast.grade)569 570    def _grade_layout(self) -> float:571        return self._safe_grade(layout.grade)572 573    def _grade_cleanliness(self) -> float:574        return self._safe_grade(cleanliness.grade)575 576    def _grade_design_quality(self) -> float:577        return self._safe_grade(design_quality.grade)578 579    def _is_done(self, scores: Dict[str, float]) -> bool:580        required = list(self.state_data.get("required_graders") or self.required_graders)581        if not required:582            required = ["color", "spacing", "typography", "contrast", "cleanliness"]583        return all(scores.get(key, 0.0) >= self.success_threshold for key in required)584 585    def _safe_grade(self, grader_fn) -> float:586        try:587            score = float(grader_fn(self.html, self.css, self.tokens, self.state_data))588            return self._clamp_open_unit_interval(score)589        except Exception:590            return self.MIN_SCORE_BOUND + self.SCORE_EPSILON591 592    def _compute_unused_selectors(self, css: str, html: str) -> List[str]:593        selector_groups = re.findall(r"([^{}]+)\{", css)594        selectors: List[str] = []595 596        for group in selector_groups:597            for selector in group.split(","):598                clean = selector.strip()599                if clean and clean not in selectors:600                    selectors.append(clean)601 602        unused = [s for s in selectors if not self._selector_matches_html(s, html)]603        return sorted(set(unused))604 605    def _selector_matches_html(self, selector: str, html: str) -> bool:606        # Ignore pseudo and combinator suffixes and evaluate only the right-most simple token.607        base = re.split(r"::?|\s+|>|\+|~", selector.strip())[-1]608        if not base or base == "*":609            return True610 611        if base.startswith("."):612            class_name = base[1:]613            return bool(re.search(rf'class\s*=\s*["\'][^"\']*\b{re.escape(class_name)}\b', html))614 615        if base.startswith("#"):616            element_id = base[1:]617            return bool(re.search(rf'id\s*=\s*["\']{re.escape(element_id)}["\']', html))618 619        if re.match(r"^[a-zA-Z][a-zA-Z0-9-]*$", base):620            return bool(re.search(rf"<{re.escape(base)}(?:\s|>)", html))621 622        return False623