openenv/echo_env
6
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 