openenv/coding_env
21
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 