Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
containers.py575 linesDownload Raw Back to rubrics
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"""Container rubrics for composing reward computations.8 9These containers provide common aggregation patterns for rubrics,10similar to how PyTorch provides nn.Sequential alongside nn.Module.11 12See RFC 004 for full design: rfcs/004-rubrics.md13"""14 15import asyncio16import inspect17from typing import Any, Dict, Iterator, List, Mapping, Tuple, Union18 19from openenv.core.rubrics.base import Rubric20 21 22def _in_async_context() -> bool:23    """Check if we're currently in an async context."""24    try:25        asyncio.get_running_loop()26        return True27    except RuntimeError:28        return False29 30 31class Sequential(Rubric):32    """Run rubrics in order, fail-fast on zero.33 34    Runs child rubrics in order. If any returns 0, stops immediately35    and returns 0. This implements hierarchical gating patterns where36    syntax checks run before execution checks.37 38    Usage:39        rubric = Sequential(40            Gate(Compiles()),41            Gate(PassesTests(), threshold=0.5),42            WeightedSum([PassesTests(), StyleRubric()], weights=[0.7, 0.3])43        )44    """45 46    def __init__(self, *rubrics: Rubric):47        """Initialize with rubrics to run in sequence.48 49        Args:50            *rubrics: Rubrics to run in order. Stops and returns 0 if any51                child returns 0.52        """53        super().__init__()54        for i, rubric in enumerate(rubrics):55            setattr(self, f"rubric_{i}", rubric)56        self._rubric_list = list(rubrics)57 58    def forward(self, action: Any, observation: Any) -> float:59        """Run rubrics in order, return 0 if any returns 0. Sync version."""60        result = 1.061        for rubric in self._rubric_list:62            score = rubric(action, observation)63            if score == 0.0:64                return 0.065            result = score66        return result67 68    def __call__(self, action: Any, observation: Any):69        """Override to choose sync or async path based on children."""70        # Empty case - check if in async context71        if not self._rubric_list:72            if _in_async_context():73                return self._empty_async(action, observation)74            else:75                # Pre-hooks76                for hook in self._forward_pre_hooks:77                    hook(self, action, observation)78                result = 1.079                self.last_score = result80                for hook in self._forward_hooks:81                    hook(self, action, observation, result)82                return result83 84        # Call first rubric to see if it's async85        first_result = self._rubric_list[0](action, observation)86        if inspect.iscoroutine(first_result):87            # At least one child is async, use async path88            return self._call_async_detected(action, observation, first_result)89        else:90            # Continue with sync path91            if first_result == 0.0:92                # Pre-hooks93                for hook in self._forward_pre_hooks:94                    hook(self, action, observation)95                self.last_score = 0.096                for hook in self._forward_hooks:97                    hook(self, action, observation, 0.0)98                return 0.099 100            final_result = first_result101            for i, rubric in enumerate(self._rubric_list[1:], start=1):102                score = rubric(action, observation)103                if inspect.iscoroutine(score):104                    # Found async mid-way, switch to async105                    # We already called rubric at index i, so pass the coroutine and remaining rubrics106                    return self._call_async_mid(107                        action,108                        observation,109                        final_result,110                        score,111                        self._rubric_list[i + 1 :],112                    )113                if score == 0.0:114                    # Pre-hooks115                    for hook in self._forward_pre_hooks:116                        hook(self, action, observation)117                    self.last_score = 0.0118                    for hook in self._forward_hooks:119                        hook(self, action, observation, 0.0)120                    return 0.0121                final_result = score122 123            # All sync - check if in async context124            if _in_async_context():125                return self._wrap_sync_result(action, observation, final_result)126            else:127                # Pre-hooks128                for hook in self._forward_pre_hooks:129                    hook(self, action, observation)130                self.last_score = final_result131                for hook in self._forward_hooks:132                    hook(self, action, observation, final_result)133                return final_result134 135    async def _empty_async(self, action, observation):136        """Async path for empty sequential."""137        for hook in self._forward_pre_hooks:138            if inspect.iscoroutinefunction(hook):139                await hook(self, action, observation)140            else:141                hook(self, action, observation)142 143        result = 1.0144        self.last_score = result145 146        for hook in self._forward_hooks:147            if inspect.iscoroutinefunction(hook):148                await hook(self, action, observation, result)149            else:150                hook(self, action, observation, result)151        return result152 153    async def _wrap_sync_result(self, action, observation, result):154        """Wrap sync result for async context."""155        for hook in self._forward_pre_hooks:156            if inspect.iscoroutinefunction(hook):157                await hook(self, action, observation)158            else:159                hook(self, action, observation)160 161        self.last_score = result162 163        for hook in self._forward_hooks:164            if inspect.iscoroutinefunction(hook):165                await hook(self, action, observation, result)166            else:167                hook(self, action, observation, result)168        return result169 170    async def _call_async_detected(self, action, observation, first_coro):171        """Async path when first child is async."""172        for hook in self._forward_pre_hooks:173            if inspect.iscoroutinefunction(hook):174                await hook(self, action, observation)175            else:176                hook(self, action, observation)177 178        result = await first_coro179        if result == 0.0:180            self.last_score = 0.0181            for hook in self._forward_hooks:182                if inspect.iscoroutinefunction(hook):183                    await hook(self, action, observation, result)184                else:185                    hook(self, action, observation, result)186            return 0.0187 188        for rubric in self._rubric_list[1:]:189            score = rubric(action, observation)190            if inspect.iscoroutine(score):191                score = await score192            if score == 0.0:193                self.last_score = 0.0194                for hook in self._forward_hooks:195                    if inspect.iscoroutinefunction(hook):196                        await hook(self, action, observation, 0.0)197                    else:198                        hook(self, action, observation, 0.0)199                return 0.0200            result = score201 202        self.last_score = result203        for hook in self._forward_hooks:204            if inspect.iscoroutinefunction(hook):205                await hook(self, action, observation, result)206            else:207                hook(self, action, observation, result)208        return result209 210    async def _call_async_mid(211        self, action, observation, current_result, first_async_coro, remaining212    ):213        """Async path when async detected mid-execution."""214        for hook in self._forward_pre_hooks:215            if inspect.iscoroutinefunction(hook):216                await hook(self, action, observation)217            else:218                hook(self, action, observation)219 220        # Await the first async rubric (already called)221        result = await first_async_coro222        if result == 0.0:223            self.last_score = 0.0224            for hook in self._forward_hooks:225                if inspect.iscoroutinefunction(hook):226                    await hook(self, action, observation, 0.0)227                else:228                    hook(self, action, observation, 0.0)229            return 0.0230 231        # Continue with remaining rubrics232        for rubric in remaining:233            score = rubric(action, observation)234            if inspect.iscoroutine(score):235                score = await score236            if score == 0.0:237                self.last_score = 0.0238                for hook in self._forward_hooks:239                    if inspect.iscoroutinefunction(hook):240                        await hook(self, action, observation, 0.0)241                    else:242                        hook(self, action, observation, 0.0)243                return 0.0244            result = score245 246        self.last_score = result247        for hook in self._forward_hooks:248            if inspect.iscoroutinefunction(hook):249                await hook(self, action, observation, result)250            else:251                hook(self, action, observation, result)252        return result253 254    def __len__(self) -> int:255        return len(self._rubric_list)256 257    def __getitem__(self, index: int) -> Rubric:258        return self._rubric_list[index]259 260 261class Gate(Rubric):262    """Threshold wrapper - returns 0 if child score is below threshold.263 264    Useful for hard constraints like "must pass 50% of tests".265 266    Usage:267        rubric = Gate(PassesTests(), threshold=0.5)268        # Returns PassesTests() score if >= 0.5, else 0.0269    """270 271    def __init__(self, rubric: Rubric, threshold: float = 1.0):272        """Initialize with a rubric and threshold.273 274        Args:275            rubric: The rubric to gate.276            threshold: Minimum score required. If child returns less than277                this, Gate returns 0. Default is 1.0 (must pass completely).278        """279        super().__init__()280        self.rubric = rubric281        self.threshold = threshold282 283    def forward(self, action: Any, observation: Any) -> float:284        """Return child score if >= threshold, else 0. Sync version."""285        score = self.rubric(action, observation)286        if score < self.threshold:287            return 0.0288        return score289 290    def __call__(self, action: Any, observation: Any):291        """Override to handle async child."""292        # Call child293        score = self.rubric(action, observation)294 295        if inspect.iscoroutine(score):296            # Child is async297            return self._call_async(action, observation, score)298        else:299            # Child is sync300            # Pre-hooks301            for hook in self._forward_pre_hooks:302                hook(self, action, observation)303            result = 0.0 if score < self.threshold else score304            self.last_score = result305            for hook in self._forward_hooks:306                hook(self, action, observation, result)307            return result308 309    async def _call_async(self, action, observation, score_coro):310        """Async path."""311        for hook in self._forward_pre_hooks:312            if inspect.iscoroutinefunction(hook):313                await hook(self, action, observation)314            else:315                hook(self, action, observation)316 317        score = await score_coro318        result = 0.0 if score < self.threshold else score319        self.last_score = result320 321        for hook in self._forward_hooks:322            if inspect.iscoroutinefunction(hook):323                await hook(self, action, observation, result)324            else:325                hook(self, action, observation, result)326        return result327 328 329class WeightedSum(Rubric):330    """Weighted combination of child rubrics.331 332    Standard aggregation pattern for multi-criteria evaluation.333 334    Usage:335        rubric = WeightedSum(336            [PassesTests(), StyleRubric()],337            weights=[0.7, 0.3]338        )339    """340 341    def __init__(self, rubrics: List[Rubric], weights: List[float]):342        """Initialize with rubrics and weights.343 344        Args:345            rubrics: List of rubrics to combine.346            weights: Weight for each rubric. Must sum to 1.0.347 348        Raises:349            ValueError: If lengths don't match or weights don't sum to 1.0.350        """351        super().__init__()352        if len(rubrics) != len(weights):353            raise ValueError(354                f"Number of rubrics ({len(rubrics)}) must match "355                f"number of weights ({len(weights)})"356            )357        if abs(sum(weights) - 1.0) > 1e-6:358            raise ValueError(f"Weights must sum to 1.0, got {sum(weights)}")359 360        for i, rubric in enumerate(rubrics):361            setattr(self, f"rubric_{i}", rubric)362        self._rubric_list = list(rubrics)363        self._weights = list(weights)364 365    def forward(self, action: Any, observation: Any) -> float:366        """Return weighted sum of child scores. Sync version."""367        total = 0.0368        for rubric, weight in zip(self._rubric_list, self._weights):369            score = rubric(action, observation)370            total += score * weight371        return total372 373    def __call__(self, action: Any, observation: Any):374        """Override to handle async children with parallel execution."""375        # Call all rubrics376        results = [rubric(action, observation) for rubric in self._rubric_list]377 378        # Check if any are async379        has_async = any(inspect.iscoroutine(r) for r in results)380 381        if has_async:382            # Use async path383            return self._call_async(action, observation, results)384        else:385            # Sync path386            # Pre-hooks387            for hook in self._forward_pre_hooks:388                hook(self, action, observation)389            total = 0.0390            for score, weight in zip(results, self._weights):391                total += score * weight392            self.last_score = total393            for hook in self._forward_hooks:394                hook(self, action, observation, total)395            return total396 397    async def _call_async(self, action, observation, results):398        """Async path with parallel execution."""399        for hook in self._forward_pre_hooks:400            if inspect.iscoroutinefunction(hook):401                await hook(self, action, observation)402            else:403                hook(self, action, observation)404 405        # Separate sync and async results406        async_tasks = []407        async_indices = []408        scores = [None] * len(results)409 410        for i, result in enumerate(results):411            if inspect.iscoroutine(result):412                async_tasks.append(result)413                async_indices.append(i)414            else:415                scores[i] = result416 417        # Await all async tasks in parallel418        if async_tasks:419            async_scores = await asyncio.gather(*async_tasks)420            for i, score in zip(async_indices, async_scores):421                scores[i] = score422 423        # Compute weighted sum424        total = 0.0425        for score, weight in zip(scores, self._weights):426            total += score * weight427 428        self.last_score = total429 430        for hook in self._forward_hooks:431            if inspect.iscoroutinefunction(hook):432                await hook(self, action, observation, total)433            else:434                hook(self, action, observation, total)435        return total436 437    @property438    def weights(self) -> List[float]:439        """Get the weights (read-only copy)."""440        return list(self._weights)441 442 443class RubricList(Rubric):444    """Container for dynamic lists of rubrics.445 446    Analogous to nn.ModuleList. Does not define aggregation - use within447    a parent rubric that implements custom logic.448 449    Usage:450        class MultiGameRubric(Rubric):451            def __init__(self, games: List[str]):452                super().__init__()453                self.games = RubricList([GameRubric(g) for g in games])454 455            def forward(self, action, obs) -> float:456                return self.games[obs.game_index](action, obs)457    """458 459    def __init__(self, rubrics: List[Rubric] = None):460        """Initialize with optional list of rubrics.461 462        Args:463            rubrics: Optional list of rubrics to start with.464        """465        super().__init__()466        self._rubrics: List[Rubric] = []467        if rubrics is not None:468            for i, rubric in enumerate(rubrics):469                self.append(rubric)470 471    def forward(self, action: Any, observation: Any) -> float:472        """RubricList does not define aggregation - override in parent."""473        raise NotImplementedError(474            "RubricList.forward() is not implemented. "475            "Use RubricList within a parent rubric that defines aggregation."476        )477 478    def append(self, rubric: Rubric) -> None:479        """Add a rubric to the list."""480        index = len(self._rubrics)481        setattr(self, f"rubric_{index}", rubric)482        self._rubrics.append(rubric)483 484    def extend(self, rubrics: List[Rubric]) -> None:485        """Add multiple rubrics to the list."""486        for rubric in rubrics:487            self.append(rubric)488 489    def __len__(self) -> int:490        return len(self._rubrics)491 492    def __getitem__(self, index: int) -> Rubric:493        return self._rubrics[index]494 495    def __iter__(self) -> Iterator[Rubric]:496        return iter(self._rubrics)497 498 499class RubricDict(Rubric):500    """Container for named rubrics with keyed access.501 502    Analogous to nn.ModuleDict. Enables keyed access for multi-task503    environments where different tasks require different rubrics.504 505    Usage:506        class AtariRubric(Rubric):507            def __init__(self):508                super().__init__()509                self.games = RubricDict({510                    "pong": PongRubric(),511                    "breakout": BreakoutRubric(),512                    "space_invaders": SpaceInvadersRubric(),513                })514 515            def forward(self, action, obs) -> float:516                return self.games[obs.game_id](action, obs)517 518        # Access: env.rubric.games["pong"]519    """520 521    def __init__(self, rubrics: Dict[str, Rubric] = None):522        """Initialize with optional dictionary of rubrics.523 524        Args:525            rubrics: Optional dictionary mapping names to rubrics.526        """527        super().__init__()528        self._rubric_dict: Dict[str, Rubric] = {}529        if rubrics is not None:530            for name, rubric in rubrics.items():531                self[name] = rubric532 533    def forward(self, action: Any, observation: Any) -> float:534        """RubricDict does not define aggregation - override in parent."""535        raise NotImplementedError(536            "RubricDict.forward() is not implemented. "537            "Use RubricDict within a parent rubric that defines aggregation."538        )539 540    def __setitem__(self, key: str, rubric: Rubric) -> None:541        """Add a rubric with the given key."""542        setattr(self, key, rubric)543        self._rubric_dict[key] = rubric544 545    def __getitem__(self, key: str) -> Rubric:546        """Get rubric by key."""547        return self._rubric_dict[key]548 549    def __contains__(self, key: str) -> bool:550        """Check if key exists."""551        return key in self._rubric_dict552 553    def __len__(self) -> int:554        return len(self._rubric_dict)555 556    def __iter__(self) -> Iterator[str]:557        return iter(self._rubric_dict)558 559    def keys(self) -> Iterator[str]:560        """Iterate over keys."""561        return iter(self._rubric_dict.keys())562 563    def values(self) -> Iterator[Rubric]:564        """Iterate over rubrics."""565        return iter(self._rubric_dict.values())566 567    def items(self) -> Iterator[Tuple[str, Rubric]]:568        """Iterate over (key, rubric) pairs."""569        return iter(self._rubric_dict.items())570 571    def update(self, rubrics: Union[Dict[str, Rubric], Mapping[str, Rubric]]) -> None:572        """Update with rubrics from a dictionary."""573        for name, rubric in rubrics.items():574            self[name] = rubric575