Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 2d agoView on Hugging Face
6likes
base.py217 linesDownload Raw Back to rubrics
1# SPDX-License-Identifier: BSD-3-Clause2 3"""Base Rubric class for reward computation.4 5Rubrics compute rewards from actions and observations. The API is modeled6after PyTorch's nn.Module: users implement forward(), and the framework7handles child registration and hooks.8 9See RFC 004 for full design: rfcs/004-rubrics.md10"""11 12import inspect13from abc import ABC, abstractmethod14from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple15 16 17class Rubric(ABC):18    """Abstract base class for reward computation.19 20    A Rubric computes a reward signal from an action and observation.21    Subclasses implement forward() to define the reward logic.22 23    Examples:24 25        ```python26        class MyRubric(Rubric):27            def forward(self, action, observation) -> float:28                return 1.0 if action.valid else 0.029 30        rubric = MyRubric()31        reward = rubric(action, observation)32        ```33 34    Child rubrics are auto-registered when assigned as attributes,35    enabling hierarchical composition and introspection.36    """37 38    _rubric_children: Dict[str, "Rubric"]39    _forward_hooks: List[Callable]40    _forward_pre_hooks: List[Callable]41    last_score: Optional[float]42 43    def __init__(self):44        # Use object.__setattr__ to avoid triggering __setattr__ during init45        object.__setattr__(self, "_rubric_children", {})46        object.__setattr__(self, "_forward_hooks", [])47        object.__setattr__(self, "_forward_pre_hooks", [])48        object.__setattr__(self, "last_score", None)49 50    def __setattr__(self, name: str, value: Any) -> None:51        # Auto-register child rubrics when assigned as attributes52        if isinstance(value, Rubric):53            self._rubric_children[name] = value54        object.__setattr__(self, name, value)55 56    def __call__(self, action: Any, observation: Any):57        """Evaluate the rubric with hooks.58 59        Args:60            action: The action taken by the agent.61            observation: The resulting observation.62 63        Returns:64            `float`: Reward value (typically 0.0 to 1.0).65        """66        # Check if forward method is async BEFORE calling it67        if inspect.iscoroutinefunction(self.forward):68            # Async path - pre-hooks will be called in _call_async69            result = self.forward(action, observation)70            return self._call_async(action, observation, result)71        else:72            # Sync path - call pre-hooks BEFORE forward()73            for hook in self._forward_pre_hooks:74                hook(self, action, observation)75            result = self.forward(action, observation)76            return self._call_sync(action, observation, result)77 78    def _call_sync(self, action: Any, observation: Any, result: float) -> float:79        """Synchronous call path."""80        return self._finish_forward(action, observation, result)81 82    def _run_forward_pre_hooks(self, action: Any, observation: Any) -> None:83        """Run pre-forward hooks synchronously."""84        for hook in self._forward_pre_hooks:85            hook(self, action, observation)86 87    async def _run_forward_pre_hooks_async(self, action: Any, observation: Any) -> None:88        """Run pre-forward hooks from an async call path."""89        for hook in self._forward_pre_hooks:90            if inspect.iscoroutinefunction(hook):91                await hook(self, action, observation)92            else:93                hook(self, action, observation)94 95    def _finish_forward(self, action: Any, observation: Any, result: float) -> float:96        """Store the result and run post-forward hooks synchronously."""97        self.last_score = result98 99        # Post-forward hooks100        for hook in self._forward_hooks:101            hook(self, action, observation, result)102 103        return result104 105    async def _finish_forward_async(106        self, action: Any, observation: Any, result: float107    ) -> float:108        """Store the result and run post-forward hooks from an async call path."""109        self.last_score = result110 111        # Post-forward hooks112        for hook in self._forward_hooks:113            if inspect.iscoroutinefunction(hook):114                await hook(self, action, observation, result)115            else:116                hook(self, action, observation, result)117 118        return result119 120    async def _call_async(self, action: Any, observation: Any, result_coro) -> float:121        """Asynchronous call path."""122        # Pre-forward hooks123        await self._run_forward_pre_hooks_async(action, observation)124 125        # Await the forward result126        result = await result_coro127        return await self._finish_forward_async(action, observation, result)128 129    @abstractmethod130    def forward(self, action: Any, observation: Any) -> float:131        """Compute the reward. Implement this in subclasses.132 133        Args:134            action: The action taken by the agent.135            observation: The resulting observation.136 137        Returns:138            `float`: Reward value (typically 0.0 to 1.0).139        """140        raise NotImplementedError141 142    def register_forward_hook(143        self, hook: Callable[["Rubric", Any, Any, float], None]144    ) -> None:145        """Register a hook called after forward().146 147        Args:148            hook (`Callable`):149                Callable with signature (rubric, action, observation, result).150        """151        self._forward_hooks.append(hook)152 153    def register_forward_pre_hook(154        self, hook: Callable[["Rubric", Any, Any], None]155    ) -> None:156        """Register a hook called before forward().157 158        Args:159            hook (`Callable`):160                Callable with signature (rubric, action, observation).161        """162        self._forward_pre_hooks.append(hook)163 164    def children(self) -> Iterator["Rubric"]:165        """Iterate over immediate child rubrics."""166        yield from self._rubric_children.values()167 168    def named_children(self) -> Iterator[Tuple[str, "Rubric"]]:169        """Iterate over immediate child rubrics with names."""170        yield from self._rubric_children.items()171 172    def rubrics(self) -> Iterator["Rubric"]:173        """Iterate over all descendant rubrics (depth-first)."""174        for child in self._rubric_children.values():175            yield child176            yield from child.rubrics()177 178    def named_rubrics(self, prefix: str = "") -> Iterator[Tuple[str, "Rubric"]]:179        """Iterate over all descendant rubrics with dot-separated names."""180        for name, child in self._rubric_children.items():181            full_name = f"{prefix}.{name}" if prefix else name182            yield full_name, child183            yield from child.named_rubrics(full_name)184 185    def get_rubric(self, path: str) -> "Rubric":186        """Access a nested rubric by dot-separated path.187 188        Args:189            path (`str`):190                Dot-separated path (e.g., "code.syntax").191 192        Returns:193            `Rubric`: The rubric at the specified path.194 195        Raises:196            KeyError: If the path does not exist.197        """198        parts = path.split(".")199        current = self200        for part in parts:201            if part not in current._rubric_children:202                raise KeyError(f"Rubric path not found: {path}")203            current = current._rubric_children[part]204        return current205 206    def reset(self) -> None:207        """Reset any internal state. Override in subclasses if needed."""208        pass209 210    def state_dict(self) -> Dict[str, Any]:211        """Serialize rubric configuration for checkpointing."""212        return {}213 214    def load_state_dict(self, state: Dict[str, Any]) -> None:215        """Load rubric configuration from checkpoint."""216        pass217