1
Fork 0
This repository has been archived on 2026-08-15. You can view files and clone it, but you cannot make any changes to it's state, such as pushing and creating new issues, pull requests or comments.
2026SEL3-project-Brittle_St.../src/brittle_star_project/evaluation/policy.py

177 lines
6.1 KiB
Python

from __future__ import annotations
from pathlib import Path
from typing import Any, Protocol
import jax
import jax.numpy as jnp
import numpy as np
from brittle_star_project.evaluation.checkpoint import load_params
class ControlPolicy(Protocol):
"""Protocol for any policy that can produce actions from observations."""
def act(self, *, observations: dict[str, Any]) -> np.ndarray: ...
class PolicyAgent:
"""Wraps a trained Flax actor for deterministic inference."""
def __init__(
self,
*,
sensor_params: Any,
actor_params: Any,
message_passer_params: Any | None = None,
message_passing_steps: int | None = None,
adj_matrix: Any | None = None,
action_dim: int,
obs_processor: Any,
) -> None:
from brittle_star_project.MLPs.mlps import (
Actor,
GenericDenseLayersWithActivation,
MessagePasser,
)
# Infer layer sizes from params
try:
dense_params = (
sensor_params.get("params", {})
if isinstance(sensor_params, dict)
else sensor_params["params"]
)
except Exception:
dense_params = sensor_params
layer_sizes = []
idx = 0
while True:
key = f"Dense_{idx}"
if key not in dense_params:
break
layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[-1]))
idx += 1
if not layer_sizes:
raise ValueError("Could not infer Dense_* layers from sensor params")
self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
self._actor = Actor(action_dim=action_dim)
self._message_passer = None
if message_passer_params is not None and not (
isinstance(message_passer_params, dict) and len(message_passer_params) == 0
):
if message_passing_steps is None or adj_matrix is None:
raise ValueError(
"Checkpoint contains message_passer_params but PolicyAgent was not given "
"message_passing_steps and adj_matrix. Pass these when constructing the agent "
"so decentralized evaluation matches training."
)
hidden_dim = int(layer_sizes[-1])
self._message_passer = MessagePasser(
hidden_dim=hidden_dim,
num_propagation_steps=int(message_passing_steps),
adj_matrix=jnp.asarray(adj_matrix),
)
self._message_passer.apply = jax.jit(self._message_passer.apply)
self._sensor.apply = jax.jit(self._sensor.apply)
self._actor.apply = jax.jit(self._actor.apply)
self._params = {
"sensor_params": sensor_params,
"actor_params": actor_params,
"message_passer_params": message_passer_params,
}
self._obs_processor = obs_processor
@classmethod
def from_params(
cls,
*,
sensor_params: Any,
actor_params: Any,
message_passer_params: Any | None = None,
message_passing_steps: int | None = None,
adj_matrix: Any | None = None,
action_dim: int,
obs_processor: Any,
) -> "PolicyAgent":
"""Construct a PolicyAgent directly from in-memory parameters."""
return cls(
sensor_params=sensor_params,
actor_params=actor_params,
message_passer_params=message_passer_params,
message_passing_steps=message_passing_steps,
adj_matrix=adj_matrix,
action_dim=action_dim,
obs_processor=obs_processor,
)
def set_params(
self,
*,
sensor_params: Any,
actor_params: Any,
message_passer_params: Any | None = None,
) -> None:
"""Update parameters for evaluation without rebuilding the model."""
self._params["sensor_params"] = sensor_params
self._params["actor_params"] = actor_params
self._params["message_passer_params"] = message_passer_params
@classmethod
def from_checkpoint(
cls,
model_path: Path,
*,
action_dim: int,
obs_processor: Any,
message_passing_steps: int | None = None,
adj_matrix: Any | None = None,
) -> "PolicyAgent":
"""Load params from .flax and construct the agent."""
params = load_params(model_path)
return cls(
sensor_params=params["sensor_params"],
actor_params=params["actor_params"],
message_passer_params=params.get("message_passer_params"),
message_passing_steps=message_passing_steps,
adj_matrix=adj_matrix,
action_dim=action_dim,
obs_processor=obs_processor,
)
def _apply_per_node(self, net, params, x):
# params: (nodes, ...)
# x: (batch, nodes, feat)
def apply_single_node(p, x_node):
# x_node: (batch, feat)
return jax.vmap(lambda xi: net.apply(p, xi))(x_node)
return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x)
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
"""Return deterministic action (actor mean, no exploration noise)."""
batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations)
obs = self._obs_processor(batched_obs)
hidden = self._apply_per_node(self._sensor, self._params["sensor_params"], obs)
if self._message_passer is not None:
mp_params = self._params.get("message_passer_params")
if mp_params is None or (isinstance(mp_params, dict) and len(mp_params) == 0):
raise ValueError(
"PolicyAgent has a message passer but message_passer_params are missing/empty."
)
hidden = jax.vmap(lambda x: self._message_passer.apply(mp_params, x))(hidden)
mean, _log_std = self._apply_per_node(self._actor, self._params["actor_params"], hidden)
return np.asarray(mean, dtype=np.float32).ravel()