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"""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 