Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
base.py196 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"""Base Rubric class for reward computation.8 9Rubrics compute rewards from actions and observations. The API is modeled10after PyTorch's nn.Module: users implement forward(), and the framework11handles child registration and hooks.12 13See RFC 004 for full design: rfcs/004-rubrics.md14"""15 16import inspect17from abc import ABC, abstractmethod18from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple19 20 21class Rubric(ABC):22    """Abstract base class for reward computation.23 24    A Rubric computes a reward signal from an action and observation.25    Subclasses implement forward() to define the reward logic.26 27    Usage:28        class MyRubric(Rubric):29            def forward(self, action, observation) -> float:30                return 1.0 if action.valid else 0.031 32        rubric = MyRubric()33        reward = rubric(action, observation)34 35    Child rubrics are auto-registered when assigned as attributes,36    enabling hierarchical composition and introspection.37    """38 39    _rubric_children: Dict[str, "Rubric"]40    _forward_hooks: List[Callable]41    _forward_pre_hooks: List[Callable]42    last_score: Optional[float]43 44    def __init__(self):45        # Use object.__setattr__ to avoid triggering __setattr__ during init46        object.__setattr__(self, "_rubric_children", {})47        object.__setattr__(self, "_forward_hooks", [])48        object.__setattr__(self, "_forward_pre_hooks", [])49        object.__setattr__(self, "last_score", None)50 51    def __setattr__(self, name: str, value: Any) -> None:52        # Auto-register child rubrics when assigned as attributes53        if isinstance(value, Rubric):54            self._rubric_children[name] = value55        object.__setattr__(self, name, value)56 57    def __call__(self, action: Any, observation: Any):58        """Evaluate the rubric with hooks.59 60        Args:61            action: The action taken by the agent.62            observation: The resulting observation.63 64        Returns:65            Reward value (typically 0.0 to 1.0).66        """67        # Check if forward method is async BEFORE calling it68        if inspect.iscoroutinefunction(self.forward):69            # Async path - pre-hooks will be called in _call_async70            result = self.forward(action, observation)71            return self._call_async(action, observation, result)72        else:73            # Sync path - call pre-hooks BEFORE forward()74            for hook in self._forward_pre_hooks:75                hook(self, action, observation)76            result = self.forward(action, observation)77            return self._call_sync(action, observation, result)78 79    def _call_sync(self, action: Any, observation: Any, result: float) -> float:80        """Synchronous call path."""81        self.last_score = result82 83        # Post-forward hooks84        for hook in self._forward_hooks:85            hook(self, action, observation, result)86 87        return result88 89    async def _call_async(self, action: Any, observation: Any, result_coro) -> float:90        """Asynchronous call path."""91        # Pre-forward hooks92        for hook in self._forward_pre_hooks:93            if inspect.iscoroutinefunction(hook):94                await hook(self, action, observation)95            else:96                hook(self, action, observation)97 98        # Await the forward result99        result = await result_coro100        self.last_score = result101 102        # Post-forward hooks103        for hook in self._forward_hooks:104            if inspect.iscoroutinefunction(hook):105                await hook(self, action, observation, result)106            else:107                hook(self, action, observation, result)108 109        return result110 111    @abstractmethod112    def forward(self, action: Any, observation: Any) -> float:113        """Compute the reward. Implement this in subclasses.114 115        Args:116            action: The action taken by the agent.117            observation: The resulting observation.118 119        Returns:120            Reward value (typically 0.0 to 1.0).121        """122        raise NotImplementedError123 124    def register_forward_hook(125        self, hook: Callable[["Rubric", Any, Any, float], None]126    ) -> None:127        """Register a hook called after forward().128 129        Args:130            hook: Callable with signature (rubric, action, observation, result).131        """132        self._forward_hooks.append(hook)133 134    def register_forward_pre_hook(135        self, hook: Callable[["Rubric", Any, Any], None]136    ) -> None:137        """Register a hook called before forward().138 139        Args:140            hook: Callable with signature (rubric, action, observation).141        """142        self._forward_pre_hooks.append(hook)143 144    def children(self) -> Iterator["Rubric"]:145        """Iterate over immediate child rubrics."""146        yield from self._rubric_children.values()147 148    def named_children(self) -> Iterator[Tuple[str, "Rubric"]]:149        """Iterate over immediate child rubrics with names."""150        yield from self._rubric_children.items()151 152    def rubrics(self) -> Iterator["Rubric"]:153        """Iterate over all descendant rubrics (depth-first)."""154        for child in self._rubric_children.values():155            yield child156            yield from child.rubrics()157 158    def named_rubrics(self, prefix: str = "") -> Iterator[Tuple[str, "Rubric"]]:159        """Iterate over all descendant rubrics with dot-separated names."""160        for name, child in self._rubric_children.items():161            full_name = f"{prefix}.{name}" if prefix else name162            yield full_name, child163            yield from child.named_rubrics(full_name)164 165    def get_rubric(self, path: str) -> "Rubric":166        """Access a nested rubric by dot-separated path.167 168        Args:169            path: Dot-separated path (e.g., "code.syntax").170 171        Returns:172            The rubric at the specified path.173 174        Raises:175            KeyError: If the path does not exist.176        """177        parts = path.split(".")178        current = self179        for part in parts:180            if part not in current._rubric_children:181                raise KeyError(f"Rubric path not found: {path}")182            current = current._rubric_children[part]183        return current184 185    def reset(self) -> None:186        """Reset any internal state. Override in subclasses if needed."""187        pass188 189    def state_dict(self) -> Dict[str, Any]:190        """Serialize rubric configuration for checkpointing."""191        return {}192 193    def load_state_dict(self, state: Dict[str, Any]) -> None:194        """Load rubric configuration from checkpoint."""195        pass196