BDO.env / server /bdo_environment.py
Vaidhav's picture
Final reward shaping for diagnosis-driven GRPO
00d6ca3
Raw
History Blame Contribute Delete
7.52 kB
from __future__ import annotations
from uuid import uuid4
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import EnvironmentMetadata
from bdo_ai_env.models import ACTION_COSTS
from bdo_ai_env.scenarios import get_scenario
from bdo_ai_env.training import consistency_bonus
from bdo_ai_env.world import BDOWorldModel
from models import BDOAction, BDOObservation, BDOState, RewardBreakdownModel, VerifierScoresModel
class BDOEnvironment(Environment[BDOAction, BDOObservation, BDOState]):
"""OpenEnv-native BDO.ai environment."""
SUPPORTS_CONCURRENT_SESSIONS: bool = True
def __init__(self, scenario: str | None = None, seed: int | None = None):
super().__init__()
self.default_scenario_name = scenario
self.default_seed = seed
self._scenario_name = scenario
self._seed = seed
self.world = BDOWorldModel(seed=seed, scenario=get_scenario(scenario))
self._state = BDOState(
episode_id=str(uuid4()),
step_count=0,
scenario_name=scenario or "randomized",
month=self.world.world.month,
shock_month=self.world.world.shock_month,
seed=seed,
terminated=False,
last_reward=None,
)
def reset(
self,
seed: int | None = None,
episode_id: str | None = None,
**kwargs,
) -> BDOObservation:
self._reset_rubric()
scenario_name = kwargs.get("scenario", self.default_scenario_name)
if scenario_name == "randomized":
scenario_name = None
scenario = get_scenario(scenario_name)
resolved_seed = seed if seed is not None else self.default_seed
state = self.world.reset(seed=resolved_seed, scenario=scenario)
self._scenario_name = scenario_name
self._seed = resolved_seed
self._state = BDOState(
episode_id=episode_id or str(uuid4()),
step_count=0,
scenario_name=state.info["scenario_name"],
month=state.info["month"],
shock_month=state.info["shock_month"],
seed=state.info["seed"],
terminated=state.terminated,
last_reward=None,
)
return self._build_observation(
reward=None,
done=False,
metadata={
"scenario_name": state.info["scenario_name"],
"month": state.info["month"],
"shock_month": state.info["shock_month"],
"seed": state.info["seed"],
"reward_breakdown": RewardBreakdownModel().model_dump(),
"verifier_scores": VerifierScoresModel().model_dump(),
"training_reward": 0.0,
"consistency_bonus": 0.0,
"executed_actions": [],
"executed_action_objects": [],
"truncated_actions": [],
"actions_truncated": False,
"errors": [],
},
)
def step(
self,
action: BDOAction,
timeout_s: float | None = None,
**kwargs,
) -> BDOObservation:
del timeout_s, kwargs
self._state.step_count += 1
syntactic_score = 0.0
info = self._apply_batch(action)
if info["actions_truncated"] or info["errors"]:
syntactic_score = -1.0
state = self.world.state()
outcome_score = (
(0.12 * state.reward.coverage)
+ (0.18 * state.reward.fraud_prevention)
+ (0.08 * state.reward.solvency)
+ (0.34 * state.reward.belief_accuracy)
)
verifier_scores = VerifierScoresModel(
syntactic_verifier=round(syntactic_score, 4),
rationality_verifier=round(state.reward.belief_accuracy, 4),
outcome_verifier=round(outcome_score, 4),
)
training_reward = round(state.reward.total + info["consistency_bonus"], 4)
reward_breakdown = RewardBreakdownModel(**state.reward.to_dict())
self._state = BDOState(
episode_id=self._state.episode_id,
step_count=self._state.step_count,
scenario_name=state.info["scenario_name"],
month=state.info["month"],
shock_month=state.info["shock_month"],
seed=state.info["seed"],
terminated=state.terminated,
last_reward=state.reward.total,
)
return self._build_observation(
reward=state.reward.total,
done=state.terminated,
metadata={
**info,
**state.info,
"reward_breakdown": reward_breakdown.model_dump(),
"verifier_scores": verifier_scores.model_dump(),
"training_reward": training_reward,
},
)
@property
def state(self) -> BDOState:
return self._state
def get_metadata(self) -> EnvironmentMetadata:
return EnvironmentMetadata(
name="BDO.ai",
description=(
"OpenEnv environment for a Block Development Officer managing "
"welfare delivery under fraud pressure, infra decay, migration, "
"and climate shocks."
),
version="0.2.0",
author="Tiya + Vaidhav",
)
def _build_observation(
self,
*,
reward: float | None,
done: bool,
metadata: dict,
) -> BDOObservation:
state = self.world.state()
payload = state.observation.to_dict()
payload["reward"] = reward
payload["done"] = done
payload["metadata"] = metadata
payload["info"] = metadata
return BDOObservation.model_validate(payload)
def _apply_batch(self, action: BDOAction) -> dict:
executed: list[str] = []
truncated: list[str] = []
errors: list[str] = []
executed_structured: list[dict[str, object]] = []
if action.predicted_fraud_level is not None:
self.world.set_predicted_fraud_level(action.predicted_fraud_level)
for command in action.actions:
cost = ACTION_COSTS[command.name]
if self.world.world.hours_remaining < cost:
truncated.append(command.name)
continue
handler = getattr(self.world, command.name, None)
if handler is None:
errors.append(f"Unknown action handler: {command.name}")
continue
try:
result = handler(**command.params)
executed.append(f"{command.name}: {result}")
executed_structured.append(
{"name": command.name, "params": dict(command.params)}
)
except Exception as exc:
errors.append(f"{command.name}: {exc}")
month_result = self.world.advance_month()
executed.append(f"advance_month: {month_result}")
bonus = consistency_bonus(action.thought_process, executed_structured)
return {
"scenario_name": self._scenario_name or "randomized",
"executed_actions": executed,
"executed_action_objects": executed_structured,
"truncated_actions": truncated,
"actions_truncated": bool(truncated),
"errors": errors,
"thought_process": action.thought_process,
"consistency_bonus": bonus,
}