Spaces:
Runtime error
Runtime error
Download policies/contextual_bandit.py from build-small-hackathon/WorldSmithAI: direct link, hf CLI and curl.
- Browser
- Download file 46.8 kB
-
https://huggingface.co/spaces/build-small-hackathon/WorldSmithAI/resolve/ee58ceded963d8b6d3657e8553a994991c97e7a3/policies/contextual_bandit.py
- Command line
-
hf download hf://spaces/build-small-hackathon/WorldSmithAI@ee58ceded963d8b6d3657e8553a994991c97e7a3/policies/contextual_bandit.py
-
curl -L -o contextual_bandit.py https://huggingface.co/spaces/build-small-hackathon/WorldSmithAI/resolve/ee58ceded963d8b6d3657e8553a994991c97e7a3/policies/contextual_bandit.py
46.8 kB
| """ | |
| Contextual bandit policy for WorldSmithAI. | |
| This module implements a domain-agnostic adaptive policy that learns which | |
| behavior to select from reward feedback. It deliberately avoids concrete | |
| behavior classes, species names, world domains, or hardcoded simulation logic. | |
| Implemented algorithms: | |
| - LinUCB: deterministic optimistic contextual bandit. | |
| - Thompson Sampling: seeded stochastic posterior sampling for reproducible | |
| experiments. | |
| Example: | |
| policy = ContextualBanditPolicy( | |
| context_keys=("state.energy", "state.credits", "memory.need_score"), | |
| alpha=1.0, | |
| ) | |
| behavior = policy.choose_action(agent, world, agent.behaviors) | |
| if behavior is not None: | |
| outcome = behavior.execute(agent, world) | |
| reward = 1.0 if outcome.get("success") else -1.0 | |
| policy.update(agent, world, behavior, reward) | |
| Future extensibility: | |
| - Add neural contextual bandits. | |
| - Add reward shaping from metrics and narrator signals. | |
| - Add off-policy evaluation. | |
| - Add per-agent shared or isolated arm states. | |
| - Add Bayesian linear regression with calibrated uncertainty. | |
| - Add exploration schedules and non-stationary environment adaptation. | |
| """ | |
| from __future__ import annotations | |
| import copy | |
| import logging | |
| from collections.abc import Iterable, Mapping, MutableMapping, MutableSequence, Sequence | |
| from dataclasses import dataclass, field | |
| from enum import Enum | |
| from numbers import Real | |
| from types import MappingProxyType | |
| from typing import TYPE_CHECKING, Any, ClassVar, cast | |
| import numpy as np | |
| from core.behavior import Behavior | |
| from policies.base_policy import BasePolicy, BehaviorCollection | |
| if TYPE_CHECKING: | |
| from core.agent import Agent | |
| from core.world import World | |
| logger = logging.getLogger(__name__) | |
| _MISSING = object() | |
| _EPSILON = 1.0e-12 | |
| class BanditAlgorithm(str, Enum): | |
| """Supported contextual bandit algorithms.""" | |
| LINUCB = "linucb" | |
| THOMPSON_SAMPLING = "thompson_sampling" | |
| class ArmKeyMode(str, Enum): | |
| """Ways to map behavior candidates to bandit arms.""" | |
| BEHAVIOR_NAME = "behavior_name" | |
| CLASS_NAME = "class_name" | |
| BEHAVIOR_NAME_WITH_INDEX = "behavior_name_with_index" | |
| class BanditEvaluation: | |
| """Scoring record for one behavior arm during action selection.""" | |
| arm_key: str | |
| behavior_name: str | |
| behavior_index: int | |
| score: float | |
| exploitation: float | |
| exploration: float | |
| uncertainty: float | |
| pulls: int | |
| context: tuple[float, ...] | |
| metadata: Mapping[str, Any] = field(default_factory=dict) | |
| def to_dict(self) -> dict[str, Any]: | |
| """Return a JSON-friendly representation of this evaluation.""" | |
| return { | |
| "arm_key": self.arm_key, | |
| "behavior_name": self.behavior_name, | |
| "behavior_index": self.behavior_index, | |
| "score": self.score, | |
| "exploitation": self.exploitation, | |
| "exploration": self.exploration, | |
| "uncertainty": self.uncertainty, | |
| "pulls": self.pulls, | |
| "context": list(self.context), | |
| "metadata": copy.deepcopy(dict(self.metadata)), | |
| } | |
| class BanditArmState: | |
| """Linear contextual bandit state for one behavior arm. | |
| For LinUCB, the state stores the precision matrix ``A`` and reward vector | |
| ``b``. The estimated parameter vector is ``theta = A^-1 b``. | |
| """ | |
| arm_key: str | |
| dimension: int | |
| ridge: float = 1.0 | |
| precision_matrix: np.ndarray = field(init=False, repr=False) | |
| reward_vector: np.ndarray = field(init=False, repr=False) | |
| pulls: int = 0 | |
| total_reward: float = 0.0 | |
| last_reward: float | None = None | |
| last_score: float | None = None | |
| last_context: tuple[float, ...] = () | |
| metadata: Mapping[str, Any] = field(default_factory=dict) | |
| def __post_init__(self) -> None: | |
| """Initialize linear model arrays.""" | |
| safe_dimension = max(1, int(self.dimension)) | |
| safe_ridge = max(_EPSILON, float(self.ridge)) | |
| self.dimension = safe_dimension | |
| self.ridge = safe_ridge | |
| self.precision_matrix = np.eye(safe_dimension, dtype=float) * safe_ridge | |
| self.reward_vector = np.zeros(safe_dimension, dtype=float) | |
| def ensure_dimension(self, dimension: int, ridge: float | None = None) -> None: | |
| """Resize model arrays when the context dimension changes. | |
| Existing values are copied into the overlapping prefix. This is useful | |
| when a DSL changes context keys between runs or during loaded-state | |
| restoration. | |
| """ | |
| safe_dimension = max(1, int(dimension)) | |
| if safe_dimension == self.dimension: | |
| return | |
| safe_ridge = max(_EPSILON, float(self.ridge if ridge is None else ridge)) | |
| old_dimension = self.dimension | |
| old_precision = self.precision_matrix | |
| old_rewards = self.reward_vector | |
| self.dimension = safe_dimension | |
| self.ridge = safe_ridge | |
| self.precision_matrix = np.eye(safe_dimension, dtype=float) * safe_ridge | |
| self.reward_vector = np.zeros(safe_dimension, dtype=float) | |
| overlap = min(old_dimension, safe_dimension) | |
| self.precision_matrix[:overlap, :overlap] = old_precision[:overlap, :overlap] | |
| self.reward_vector[:overlap] = old_rewards[:overlap] | |
| def theta(self) -> np.ndarray: | |
| """Return the estimated linear reward parameters.""" | |
| try: | |
| return np.linalg.solve(self.precision_matrix, self.reward_vector) | |
| except np.linalg.LinAlgError: | |
| logger.debug("Using pseudo-inverse for arm %s theta estimate", self.arm_key) | |
| return np.linalg.pinv(self.precision_matrix) @ self.reward_vector | |
| def inverse_precision(self) -> np.ndarray: | |
| """Return the inverse precision matrix.""" | |
| try: | |
| return np.linalg.inv(self.precision_matrix) | |
| except np.linalg.LinAlgError: | |
| logger.debug("Using pseudo-inverse for arm %s precision matrix", self.arm_key) | |
| return np.linalg.pinv(self.precision_matrix) | |
| def uncertainty(self, context: np.ndarray) -> float: | |
| """Return LinUCB uncertainty for a context vector.""" | |
| try: | |
| solved = np.linalg.solve(self.precision_matrix, context) | |
| except np.linalg.LinAlgError: | |
| solved = np.linalg.pinv(self.precision_matrix) @ context | |
| value = float(context.T @ solved) | |
| return float(np.sqrt(max(0.0, value))) | |
| def update( | |
| self, | |
| context: np.ndarray, | |
| reward: float, | |
| *, | |
| forgetting_factor: float = 1.0, | |
| ridge: float | None = None, | |
| ) -> None: | |
| """Update the linear model from one reward observation.""" | |
| self.ensure_dimension(context.size, ridge=ridge) | |
| safe_forgetting = min(max(float(forgetting_factor), 0.0), 1.0) | |
| effective_ridge = max(_EPSILON, float(self.ridge if ridge is None else ridge)) | |
| if safe_forgetting < 1.0: | |
| self.precision_matrix = ( | |
| safe_forgetting * self.precision_matrix | |
| + (1.0 - safe_forgetting) * effective_ridge * np.eye(self.dimension) | |
| ) | |
| self.reward_vector = safe_forgetting * self.reward_vector | |
| self.precision_matrix = self.precision_matrix + np.outer(context, context) | |
| self.reward_vector = self.reward_vector + float(reward) * context | |
| self.pulls += 1 | |
| self.total_reward += float(reward) | |
| self.last_reward = float(reward) | |
| self.last_context = tuple(float(value) for value in context.tolist()) | |
| def to_dict(self) -> dict[str, Any]: | |
| """Return a serializable arm-state snapshot.""" | |
| return { | |
| "arm_key": self.arm_key, | |
| "dimension": self.dimension, | |
| "ridge": self.ridge, | |
| "precision_matrix": self.precision_matrix.tolist(), | |
| "reward_vector": self.reward_vector.tolist(), | |
| "pulls": self.pulls, | |
| "total_reward": self.total_reward, | |
| "last_reward": self.last_reward, | |
| "last_score": self.last_score, | |
| "last_context": list(self.last_context), | |
| "metadata": copy.deepcopy(dict(self.metadata)), | |
| } | |
| def from_dict(cls, data: Mapping[str, Any]) -> BanditArmState: | |
| """Create arm state from a serialized mapping.""" | |
| dimension = int(data.get("dimension", 1)) | |
| arm = cls( | |
| arm_key=str(data.get("arm_key", "unknown_arm")), | |
| dimension=dimension, | |
| ridge=float(data.get("ridge", 1.0)), | |
| pulls=int(data.get("pulls", 0)), | |
| total_reward=float(data.get("total_reward", 0.0)), | |
| last_reward=None if data.get("last_reward") is None else float(data["last_reward"]), | |
| last_score=None if data.get("last_score") is None else float(data["last_score"]), | |
| last_context=tuple(float(value) for value in data.get("last_context", ())), | |
| metadata=dict(data.get("metadata", {})) | |
| if isinstance(data.get("metadata", {}), Mapping) | |
| else {}, | |
| ) | |
| precision = data.get("precision_matrix") | |
| rewards = data.get("reward_vector") | |
| if precision is not None: | |
| precision_array = np.asarray(precision, dtype=float) | |
| if precision_array.shape == (dimension, dimension): | |
| arm.precision_matrix = precision_array | |
| if rewards is not None: | |
| reward_array = np.asarray(rewards, dtype=float) | |
| if reward_array.shape == (dimension,): | |
| arm.reward_vector = reward_array | |
| return arm | |
| def _normalize_algorithm(value: BanditAlgorithm | str) -> BanditAlgorithm: | |
| """Normalize a bandit algorithm value.""" | |
| if isinstance(value, BanditAlgorithm): | |
| return value | |
| text = str(value).lower() | |
| if text in {"linucb", "lin_ucb", "ucb"}: | |
| return BanditAlgorithm.LINUCB | |
| if text in {"thompson", "thompson_sampling", "sampling"}: | |
| return BanditAlgorithm.THOMPSON_SAMPLING | |
| return BanditAlgorithm(text) | |
| def _normalize_arm_key_mode(value: ArmKeyMode | str) -> ArmKeyMode: | |
| """Normalize an arm-key mode value.""" | |
| if isinstance(value, ArmKeyMode): | |
| return value | |
| return ArmKeyMode(str(value)) | |
| def _agent_id(agent: Agent) -> str: | |
| """Return a stable string identifier for an agent.""" | |
| return str(getattr(agent, "id", "unknown_agent")) | |
| def _world_step(world: World) -> int | None: | |
| """Return the current world step if available.""" | |
| value = getattr(world, "step_count", None) | |
| if isinstance(value, Real) and not isinstance(value, bool): | |
| return int(value) | |
| return None | |
| def _is_number(value: Any) -> bool: | |
| """Return whether a value is a real numeric scalar, excluding booleans.""" | |
| return isinstance(value, (Real, np.integer, np.floating)) and not isinstance(value, bool) | |
| def _as_float(value: Any, default: float = 0.0) -> float: | |
| """Safely convert a numeric-like value to float.""" | |
| if _is_number(value): | |
| return float(value) | |
| return default | |
| def _value_as_feature(value: Any, default: float = 0.0) -> float: | |
| """Convert arbitrary context values into numeric feature values.""" | |
| if value is _MISSING or value is None: | |
| return default | |
| if isinstance(value, bool): | |
| return 1.0 if value else 0.0 | |
| if _is_number(value): | |
| return float(value) | |
| if isinstance(value, Mapping): | |
| return float(len(value)) | |
| if isinstance(value, Sequence) and not isinstance(value, (str, bytes)): | |
| return float(len(value)) | |
| return default | |
| def _split_path(path: str) -> tuple[str, ...]: | |
| """Split a dot-separated path into components.""" | |
| return tuple(part for part in str(path).split(".") if part) | |
| def _get_mapping_path(container: Mapping[str, Any], path: str, default: Any = _MISSING) -> Any: | |
| """Read a nested mapping value using dot notation.""" | |
| parts = _split_path(path) | |
| if not parts: | |
| return default | |
| current: Any = container | |
| for part in parts: | |
| if not isinstance(current, Mapping) or part not in current: | |
| return default | |
| current = current[part] | |
| return current | |
| def _get_object_path(root: Any, path: str, default: Any = _MISSING) -> Any: | |
| """Read nested values from mappings, sequences, or object attributes.""" | |
| parts = _split_path(path) | |
| if not parts: | |
| return root | |
| current: Any = root | |
| for part in parts: | |
| if isinstance(current, Mapping): | |
| if part not in current: | |
| return default | |
| current = current[part] | |
| continue | |
| if isinstance(current, Sequence) and not isinstance(current, (str, bytes)) and part.isdigit(): | |
| index = int(part) | |
| if index >= len(current): | |
| return default | |
| current = current[index] | |
| continue | |
| if not hasattr(current, part): | |
| return default | |
| current = getattr(current, part) | |
| return current | |
| def _agent_state(agent: Agent) -> Mapping[str, Any]: | |
| """Return an agent state mapping without assuming a concrete state class.""" | |
| state = getattr(agent, "state", {}) | |
| return state if isinstance(state, Mapping) else {} | |
| def _agent_memory(agent: Agent) -> Mapping[str, Any]: | |
| """Return an agent memory mapping without assuming a concrete memory class.""" | |
| memory = getattr(agent, "memory", {}) | |
| return memory if isinstance(memory, Mapping) else {} | |
| def _mutable_agent_memory(agent: Agent) -> MutableMapping[str, Any]: | |
| """Return mutable agent memory, creating one if needed.""" | |
| memory = getattr(agent, "memory", None) | |
| if isinstance(memory, MutableMapping): | |
| return memory | |
| replacement: dict[str, Any] = {} | |
| setattr(agent, "memory", replacement) | |
| return replacement | |
| def _read_context_value( | |
| agent: Agent, | |
| world: World, | |
| behavior: Behavior | None, | |
| path: str, | |
| default: Any = _MISSING, | |
| ) -> Any: | |
| """Read a generic path from agent, world, behavior, state, or memory. | |
| Supported prefixes: | |
| - ``state.foo`` | |
| - ``state:foo`` | |
| - ``memory.foo`` | |
| - ``memory:foo`` | |
| - ``agent.foo`` | |
| - ``world.foo`` | |
| - ``behavior.foo`` | |
| If no prefix is supplied, state is checked first, then memory, then agent | |
| attributes. | |
| """ | |
| normalized_path = str(path) | |
| if normalized_path.startswith("state."): | |
| return _get_mapping_path( | |
| _agent_state(agent), | |
| normalized_path.removeprefix("state."), | |
| default, | |
| ) | |
| if normalized_path.startswith("state:"): | |
| return _get_mapping_path( | |
| _agent_state(agent), | |
| normalized_path.removeprefix("state:"), | |
| default, | |
| ) | |
| if normalized_path.startswith("memory."): | |
| return _get_mapping_path( | |
| _agent_memory(agent), | |
| normalized_path.removeprefix("memory."), | |
| default, | |
| ) | |
| if normalized_path.startswith("memory:"): | |
| return _get_mapping_path( | |
| _agent_memory(agent), | |
| normalized_path.removeprefix("memory:"), | |
| default, | |
| ) | |
| if normalized_path.startswith("agent."): | |
| return _get_object_path( | |
| agent, | |
| normalized_path.removeprefix("agent."), | |
| default, | |
| ) | |
| if normalized_path.startswith("world."): | |
| return _get_object_path( | |
| world, | |
| normalized_path.removeprefix("world."), | |
| default, | |
| ) | |
| if normalized_path.startswith("behavior."): | |
| if behavior is None: | |
| return default | |
| return _get_object_path( | |
| behavior, | |
| normalized_path.removeprefix("behavior."), | |
| default, | |
| ) | |
| state_value = _get_mapping_path(_agent_state(agent), normalized_path, _MISSING) | |
| if state_value is not _MISSING: | |
| return state_value | |
| memory_value = _get_mapping_path(_agent_memory(agent), normalized_path, _MISSING) | |
| if memory_value is not _MISSING: | |
| return memory_value | |
| return _get_object_path(agent, normalized_path, default) | |
| def _append_bounded(items: MutableSequence[Any], value: Any, max_items: int) -> None: | |
| """Append an item while enforcing an optional maximum history length.""" | |
| items.append(value) | |
| if max_items > 0 and len(items) > max_items: | |
| del items[: len(items) - max_items] | |
| def _coerce_behavior_name(behavior: Behavior | str | None) -> str | None: | |
| """Return a behavior name from a behavior object or string.""" | |
| if behavior is None: | |
| return None | |
| if isinstance(behavior, str): | |
| return behavior | |
| raw_name = getattr(behavior, "name", None) | |
| if raw_name is not None: | |
| return str(raw_name) | |
| return behavior.__class__.__name__ | |
| class ContextualBanditPolicy(BasePolicy): | |
| """Adaptive contextual bandit policy for behavior selection. | |
| The policy maps behavior candidates to bandit arms and learns a linear | |
| reward model for each arm. It remains domain-agnostic by reading only | |
| generic context paths from agent state, agent memory, world fields, and | |
| behavior attributes. | |
| The default algorithm is LinUCB, which is deterministic and well-suited to | |
| reproducible simulations and hackathon demos. | |
| """ | |
| name: ClassVar[str] = "contextual_bandit" | |
| algorithm: BanditAlgorithm | str = BanditAlgorithm.LINUCB | |
| arm_key_mode: ArmKeyMode | str = ArmKeyMode.BEHAVIOR_NAME | |
| context_keys: tuple[str, ...] = () | |
| behavior_context_keys: tuple[str, ...] = () | |
| include_bias: bool = True | |
| include_bias_for_external_context: bool = True | |
| default_context_value: float = 0.0 | |
| normalize_context: bool = False | |
| max_context_norm: float | None = None | |
| clip_context_abs_value: float | None = 1.0e6 | |
| alpha: float = 1.0 | |
| thompson_variance: float = 1.0 | |
| ridge: float = 1.0 | |
| forgetting_factor: float = 1.0 | |
| new_arm_bonus: float = 0.0 | |
| reward_clip_min: float | None = None | |
| reward_clip_max: float | None = None | |
| random_seed: int | None = 0 | |
| write_decisions_to_agent_memory: bool = False | |
| decision_memory_key: str = "bandit_decisions" | |
| latest_decision_memory_key: str = "latest_bandit_decision" | |
| max_decision_history: int = 500 | |
| max_reward_history: int = 1000 | |
| _arms: dict[str, BanditArmState] = field(default_factory=dict, init=False, repr=False) | |
| _rng: np.random.Generator = field(init=False, repr=False) | |
| _decision_count: int = field(default=0, init=False, repr=False) | |
| _last_selected_arm_key: str | None = field(default=None, init=False, repr=False) | |
| _last_selected_behavior_name: str | None = field(default=None, init=False, repr=False) | |
| _last_selected_context: tuple[float, ...] = field(default_factory=tuple, init=False, repr=False) | |
| _last_evaluations: tuple[BanditEvaluation, ...] = field(default_factory=tuple, init=False, repr=False) | |
| _decision_history: list[dict[str, Any]] = field(default_factory=list, init=False, repr=False) | |
| _reward_history: list[dict[str, Any]] = field(default_factory=list, init=False, repr=False) | |
| def __post_init__(self) -> None: | |
| """Normalize configuration and initialize the random generator.""" | |
| self.algorithm = _normalize_algorithm(self.algorithm) | |
| self.arm_key_mode = _normalize_arm_key_mode(self.arm_key_mode) | |
| self.context_keys = tuple(str(key) for key in self.context_keys) | |
| self.behavior_context_keys = tuple(str(key) for key in self.behavior_context_keys) | |
| self.ridge = max(_EPSILON, float(self.ridge)) | |
| self.alpha = max(0.0, float(self.alpha)) | |
| self.thompson_variance = max(0.0, float(self.thompson_variance)) | |
| self.forgetting_factor = min(max(float(self.forgetting_factor), 0.0), 1.0) | |
| self._rng = np.random.default_rng(self.random_seed) | |
| def choose_action( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| behaviors: BehaviorCollection | None = None, | |
| ) -> Behavior | None: | |
| """Choose one behavior using the contextual bandit.""" | |
| return self.select_action(agent, world, behaviors=behaviors) | |
| def select_action( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| behaviors: BehaviorCollection | None = None, | |
| context_vector: Sequence[float] | np.ndarray | None = None, | |
| ) -> Behavior | None: | |
| """Select an action from available behaviors. | |
| Args: | |
| agent: Agent requesting an action. | |
| world: Current world. | |
| behaviors: Optional behavior collection. If omitted, reads | |
| ``agent.behaviors``. | |
| context_vector: Optional explicit context vector. When omitted, | |
| context is built from ``context_keys`` and | |
| ``behavior_context_keys``. | |
| Returns: | |
| Selected behavior, or ``None`` when no valid action exists. | |
| """ | |
| if not self.enabled: | |
| self._record_decision(agent, world, None, (), "policy_disabled") | |
| return None | |
| candidates = self.candidate_behaviors(agent, world, behaviors) | |
| valid_candidates = tuple(candidate for candidate in candidates if candidate.available) | |
| if not valid_candidates: | |
| self._record_decision(agent, world, None, (), "no_valid_behaviors") | |
| return None | |
| evaluations: list[BanditEvaluation] = [] | |
| for candidate in valid_candidates: | |
| arm_key = self.arm_key_for_behavior( | |
| candidate.behavior, | |
| behavior_name=candidate.name, | |
| behavior_index=candidate.index, | |
| ) | |
| context = self.context_vector( | |
| agent, | |
| world, | |
| behavior=candidate.behavior, | |
| explicit_context=context_vector, | |
| ) | |
| evaluation = self._evaluate_arm( | |
| arm_key=arm_key, | |
| behavior_name=candidate.name, | |
| behavior_index=candidate.index, | |
| context=context, | |
| ) | |
| evaluations.append(evaluation) | |
| selected = sorted( | |
| evaluations, | |
| key=lambda evaluation: ( | |
| -float(evaluation.score), | |
| evaluation.behavior_index, | |
| evaluation.behavior_name, | |
| evaluation.arm_key, | |
| ), | |
| )[0] | |
| selected_behavior = valid_candidates[ | |
| next( | |
| index | |
| for index, candidate in enumerate(valid_candidates) | |
| if self.arm_key_for_behavior( | |
| candidate.behavior, | |
| behavior_name=candidate.name, | |
| behavior_index=candidate.index, | |
| ) | |
| == selected.arm_key | |
| ) | |
| ].behavior | |
| self._last_selected_arm_key = selected.arm_key | |
| self._last_selected_behavior_name = selected.behavior_name | |
| self._last_selected_context = selected.context | |
| self._last_evaluations = tuple(evaluations) | |
| self.arm_state(selected.arm_key, len(selected.context)).last_score = selected.score | |
| self._record_decision(agent, world, selected, tuple(evaluations), "selected") | |
| logger.debug( | |
| "ContextualBanditPolicy selected behavior %s for agent %s with score %.3f", | |
| selected.behavior_name, | |
| _agent_id(agent), | |
| selected.score, | |
| ) | |
| return selected_behavior | |
| def update( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| behavior: Behavior | str | None, | |
| reward: float, | |
| context: Sequence[float] | Mapping[str, Any] | np.ndarray | None = None, | |
| metadata: Mapping[str, Any] | None = None, | |
| ) -> None: | |
| """Update the policy from reward feedback. | |
| Args: | |
| agent: Agent receiving feedback. | |
| world: Current world. | |
| behavior: Behavior object, behavior name, or ``None``. If ``None``, | |
| the most recently selected arm is updated. | |
| reward: Numeric reward signal. | |
| context: Optional explicit context vector or mapping. If omitted, | |
| the last selected context is reused when possible. | |
| metadata: Optional reward metadata. | |
| """ | |
| arm_key = self._arm_key_for_update(behavior) | |
| if arm_key is None: | |
| logger.debug("Skipping bandit update because no arm key is available") | |
| return | |
| context_array = self._context_for_update(agent, world, behavior, context, arm_key) | |
| if context_array is None: | |
| logger.debug("Skipping bandit update for %s because context is unavailable", arm_key) | |
| return | |
| clipped_reward = self._clip_reward(float(reward)) | |
| arm = self.arm_state(arm_key, context_array.size) | |
| arm.update( | |
| context_array, | |
| clipped_reward, | |
| forgetting_factor=self.forgetting_factor, | |
| ridge=self.ridge, | |
| ) | |
| record = { | |
| "agent_id": _agent_id(agent), | |
| "world_step": _world_step(world), | |
| "arm_key": arm_key, | |
| "behavior_name": _coerce_behavior_name(behavior) or self._last_selected_behavior_name, | |
| "reward": clipped_reward, | |
| "raw_reward": float(reward), | |
| "context": context_array.tolist(), | |
| "pulls": arm.pulls, | |
| "total_reward": arm.total_reward, | |
| "metadata": copy.deepcopy(dict(metadata or {})), | |
| } | |
| _append_bounded(self._reward_history, record, self.max_reward_history) | |
| logger.debug( | |
| "Updated bandit arm %s with reward %.3f; pulls=%s", | |
| arm_key, | |
| clipped_reward, | |
| arm.pulls, | |
| ) | |
| def receive_reward( | |
| self, | |
| reward: float, | |
| *, | |
| behavior: Behavior | str | None = None, | |
| context: Sequence[float] | Mapping[str, Any] | np.ndarray | None = None, | |
| metadata: Mapping[str, Any] | None = None, | |
| ) -> None: | |
| """Receive reward feedback when agent and world are not available. | |
| This method updates the most recently selected arm unless ``behavior`` | |
| provides a specific behavior name. | |
| """ | |
| arm_key = self._arm_key_for_update(behavior) | |
| if arm_key is None: | |
| return | |
| context_array = self._context_from_external(context) | |
| if context_array is None: | |
| if self._last_selected_context: | |
| context_array = np.asarray(self._last_selected_context, dtype=float) | |
| else: | |
| return | |
| clipped_reward = self._clip_reward(float(reward)) | |
| arm = self.arm_state(arm_key, context_array.size) | |
| arm.update( | |
| context_array, | |
| clipped_reward, | |
| forgetting_factor=self.forgetting_factor, | |
| ridge=self.ridge, | |
| ) | |
| record = { | |
| "agent_id": None, | |
| "world_step": None, | |
| "arm_key": arm_key, | |
| "behavior_name": _coerce_behavior_name(behavior) or self._last_selected_behavior_name, | |
| "reward": clipped_reward, | |
| "raw_reward": float(reward), | |
| "context": context_array.tolist(), | |
| "pulls": arm.pulls, | |
| "total_reward": arm.total_reward, | |
| "metadata": copy.deepcopy(dict(metadata or {})), | |
| } | |
| _append_bounded(self._reward_history, record, self.max_reward_history) | |
| def score_behavior( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| behavior: Behavior, | |
| ) -> float | None: | |
| """Return the current bandit score for a behavior.""" | |
| behavior_name = self.behavior_name(behavior) | |
| arm_key = self.arm_key_for_behavior(behavior, behavior_name=behavior_name) | |
| context = self.context_vector(agent, world, behavior=behavior) | |
| evaluation = self._evaluate_arm( | |
| arm_key=arm_key, | |
| behavior_name=behavior_name, | |
| behavior_index=0, | |
| context=context, | |
| ) | |
| return evaluation.score | |
| def context_vector( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| *, | |
| behavior: Behavior | None = None, | |
| explicit_context: Sequence[float] | np.ndarray | None = None, | |
| ) -> np.ndarray: | |
| """Build a numeric context vector. | |
| The context contains, in order: | |
| 1. optional bias feature | |
| 2. features from ``context_keys`` | |
| 3. behavior-specific features from ``behavior_context_keys`` | |
| Args: | |
| agent: Agent requesting action. | |
| world: Current world. | |
| behavior: Optional behavior candidate. | |
| explicit_context: Optional external context vector. If supplied, | |
| configured path keys are ignored. | |
| Returns: | |
| One-dimensional NumPy array. | |
| """ | |
| if explicit_context is not None: | |
| return self._prepare_context_array( | |
| explicit_context, | |
| include_bias=self.include_bias_for_external_context, | |
| ) | |
| values: list[float] = [] | |
| if self.include_bias: | |
| values.append(1.0) | |
| for path in self.context_keys: | |
| raw_value = _read_context_value(agent, world, behavior, path, _MISSING) | |
| values.append(_value_as_feature(raw_value, self.default_context_value)) | |
| for path in self.behavior_context_keys: | |
| raw_value = _read_context_value(agent, world, behavior, path, _MISSING) | |
| values.append(_value_as_feature(raw_value, self.default_context_value)) | |
| if not values: | |
| values.append(0.0) | |
| return self._prepare_context_array(values, include_bias=False) | |
| def arm_key_for_behavior( | |
| self, | |
| behavior: Behavior, | |
| *, | |
| behavior_name: str | None = None, | |
| behavior_index: int | None = None, | |
| ) -> str: | |
| """Return the bandit arm key for a behavior candidate.""" | |
| mode = _normalize_arm_key_mode(self.arm_key_mode) | |
| resolved_name = behavior_name or self.behavior_name(behavior) | |
| if mode is ArmKeyMode.CLASS_NAME: | |
| return behavior.__class__.__name__ | |
| if mode is ArmKeyMode.BEHAVIOR_NAME_WITH_INDEX: | |
| index = 0 if behavior_index is None else int(behavior_index) | |
| return f"{resolved_name}#{index}" | |
| return resolved_name | |
| def arm_state(self, arm_key: str, dimension: int) -> BanditArmState: | |
| """Return the mutable arm state for an arm key, creating it if needed.""" | |
| if arm_key not in self._arms: | |
| self._arms[arm_key] = BanditArmState( | |
| arm_key=arm_key, | |
| dimension=max(1, int(dimension)), | |
| ridge=self.ridge, | |
| ) | |
| arm = self._arms[arm_key] | |
| arm.ensure_dimension(max(1, int(dimension)), ridge=self.ridge) | |
| return arm | |
| def reset(self) -> None: | |
| """Reset learned arm state and runtime traces.""" | |
| self._arms.clear() | |
| self._decision_count = 0 | |
| self._last_selected_arm_key = None | |
| self._last_selected_behavior_name = None | |
| self._last_selected_context = () | |
| self._last_evaluations = () | |
| self._decision_history.clear() | |
| self._reward_history.clear() | |
| self._rng = np.random.default_rng(self.random_seed) | |
| def state_dict(self) -> dict[str, Any]: | |
| """Return a serializable policy state snapshot.""" | |
| state = super().state_dict() | |
| state.update( | |
| { | |
| "algorithm": _normalize_algorithm(self.algorithm).value, | |
| "arm_key_mode": _normalize_arm_key_mode(self.arm_key_mode).value, | |
| "context_keys": list(self.context_keys), | |
| "behavior_context_keys": list(self.behavior_context_keys), | |
| "include_bias": bool(self.include_bias), | |
| "include_bias_for_external_context": bool(self.include_bias_for_external_context), | |
| "default_context_value": float(self.default_context_value), | |
| "normalize_context": bool(self.normalize_context), | |
| "max_context_norm": self.max_context_norm, | |
| "clip_context_abs_value": self.clip_context_abs_value, | |
| "alpha": float(self.alpha), | |
| "thompson_variance": float(self.thompson_variance), | |
| "ridge": float(self.ridge), | |
| "forgetting_factor": float(self.forgetting_factor), | |
| "new_arm_bonus": float(self.new_arm_bonus), | |
| "reward_clip_min": self.reward_clip_min, | |
| "reward_clip_max": self.reward_clip_max, | |
| "random_seed": self.random_seed, | |
| "arms": {key: arm.to_dict() for key, arm in self._arms.items()}, | |
| "decision_count": int(self._decision_count), | |
| "last_selected_arm_key": self._last_selected_arm_key, | |
| "last_selected_behavior_name": self._last_selected_behavior_name, | |
| "last_selected_context": list(self._last_selected_context), | |
| "decision_history": copy.deepcopy(self._decision_history), | |
| "reward_history": copy.deepcopy(self._reward_history), | |
| "rng_state": copy.deepcopy(self._rng.bit_generator.state), | |
| } | |
| ) | |
| return state | |
| def load_state_dict(self, state: Mapping[str, Any]) -> None: | |
| """Load policy configuration and learned arm state from a mapping.""" | |
| super().load_state_dict(state) | |
| if "algorithm" in state: | |
| self.algorithm = _normalize_algorithm(str(state["algorithm"])) | |
| if "arm_key_mode" in state: | |
| self.arm_key_mode = _normalize_arm_key_mode(str(state["arm_key_mode"])) | |
| self.context_keys = tuple(str(key) for key in state.get("context_keys", self.context_keys)) | |
| self.behavior_context_keys = tuple( | |
| str(key) for key in state.get("behavior_context_keys", self.behavior_context_keys) | |
| ) | |
| if "include_bias" in state: | |
| self.include_bias = bool(state["include_bias"]) | |
| if "include_bias_for_external_context" in state: | |
| self.include_bias_for_external_context = bool(state["include_bias_for_external_context"]) | |
| if "default_context_value" in state: | |
| self.default_context_value = _as_float(state["default_context_value"], self.default_context_value) | |
| if "normalize_context" in state: | |
| self.normalize_context = bool(state["normalize_context"]) | |
| self.max_context_norm = ( | |
| None if state.get("max_context_norm") is None else float(state["max_context_norm"]) | |
| ) | |
| self.clip_context_abs_value = ( | |
| None | |
| if state.get("clip_context_abs_value") is None | |
| else float(state["clip_context_abs_value"]) | |
| ) | |
| self.alpha = max(0.0, _as_float(state.get("alpha"), self.alpha)) | |
| self.thompson_variance = max( | |
| 0.0, | |
| _as_float(state.get("thompson_variance"), self.thompson_variance), | |
| ) | |
| self.ridge = max(_EPSILON, _as_float(state.get("ridge"), self.ridge)) | |
| self.forgetting_factor = min( | |
| max(_as_float(state.get("forgetting_factor"), self.forgetting_factor), 0.0), | |
| 1.0, | |
| ) | |
| self.new_arm_bonus = _as_float(state.get("new_arm_bonus"), self.new_arm_bonus) | |
| self.reward_clip_min = ( | |
| None if state.get("reward_clip_min") is None else float(state["reward_clip_min"]) | |
| ) | |
| self.reward_clip_max = ( | |
| None if state.get("reward_clip_max") is None else float(state["reward_clip_max"]) | |
| ) | |
| self.random_seed = None if state.get("random_seed") is None else int(state["random_seed"]) | |
| self._rng = np.random.default_rng(self.random_seed) | |
| rng_state = state.get("rng_state") | |
| if isinstance(rng_state, Mapping): | |
| try: | |
| self._rng.bit_generator.state = copy.deepcopy(dict(rng_state)) | |
| except (TypeError, ValueError): | |
| logger.debug("Could not restore contextual bandit RNG state") | |
| arms = state.get("arms") | |
| if isinstance(arms, Mapping): | |
| self._arms = { | |
| str(key): BanditArmState.from_dict(value) | |
| for key, value in arms.items() | |
| if isinstance(value, Mapping) | |
| } | |
| self._decision_count = int(state.get("decision_count", self._decision_count)) | |
| self._last_selected_arm_key = ( | |
| None | |
| if state.get("last_selected_arm_key") is None | |
| else str(state["last_selected_arm_key"]) | |
| ) | |
| self._last_selected_behavior_name = ( | |
| None | |
| if state.get("last_selected_behavior_name") is None | |
| else str(state["last_selected_behavior_name"]) | |
| ) | |
| self._last_selected_context = tuple( | |
| float(value) for value in state.get("last_selected_context", ()) | |
| ) | |
| decision_history = state.get("decision_history") | |
| if isinstance(decision_history, Sequence) and not isinstance(decision_history, (str, bytes)): | |
| self._decision_history = [ | |
| copy.deepcopy(dict(item)) | |
| for item in decision_history | |
| if isinstance(item, Mapping) | |
| ] | |
| reward_history = state.get("reward_history") | |
| if isinstance(reward_history, Sequence) and not isinstance(reward_history, (str, bytes)): | |
| self._reward_history = [ | |
| copy.deepcopy(dict(item)) | |
| for item in reward_history | |
| if isinstance(item, Mapping) | |
| ] | |
| def arms(self) -> Mapping[str, BanditArmState]: | |
| """Return a read-only view of learned arm states.""" | |
| return MappingProxyType(self._arms) | |
| def last_evaluations(self) -> tuple[BanditEvaluation, ...]: | |
| """Return the latest selection evaluations.""" | |
| return self._last_evaluations | |
| def decision_history(self) -> tuple[Mapping[str, Any], ...]: | |
| """Return immutable recent decision history.""" | |
| return tuple(copy.deepcopy(self._decision_history)) | |
| def reward_history(self) -> tuple[Mapping[str, Any], ...]: | |
| """Return immutable recent reward history.""" | |
| return tuple(copy.deepcopy(self._reward_history)) | |
| def _evaluate_arm( | |
| self, | |
| *, | |
| arm_key: str, | |
| behavior_name: str, | |
| behavior_index: int, | |
| context: np.ndarray, | |
| ) -> BanditEvaluation: | |
| """Score one arm for a context vector.""" | |
| arm = self.arm_state(arm_key, context.size) | |
| theta = arm.theta() | |
| exploitation = float(context @ theta) | |
| uncertainty = arm.uncertainty(context) | |
| algorithm = _normalize_algorithm(self.algorithm) | |
| if algorithm is BanditAlgorithm.THOMPSON_SAMPLING: | |
| sampled_theta = self._sample_theta(arm) | |
| score = float(context @ sampled_theta) | |
| exploration = score - exploitation | |
| else: | |
| exploration = float(self.alpha) * uncertainty | |
| score = exploitation + exploration | |
| if arm.pulls == 0: | |
| score += float(self.new_arm_bonus) | |
| arm.last_score = score | |
| arm.last_context = tuple(float(value) for value in context.tolist()) | |
| return BanditEvaluation( | |
| arm_key=arm_key, | |
| behavior_name=behavior_name, | |
| behavior_index=behavior_index, | |
| score=score, | |
| exploitation=exploitation, | |
| exploration=exploration, | |
| uncertainty=uncertainty, | |
| pulls=arm.pulls, | |
| context=tuple(float(value) for value in context.tolist()), | |
| metadata={ | |
| "algorithm": algorithm.value, | |
| "new_arm": arm.pulls == 0, | |
| "total_reward": arm.total_reward, | |
| }, | |
| ) | |
| def _sample_theta(self, arm: BanditArmState) -> np.ndarray: | |
| """Sample arm parameters for Thompson Sampling.""" | |
| mean = arm.theta() | |
| covariance = (float(self.thompson_variance) ** 2) * arm.inverse_precision() | |
| covariance = (covariance + covariance.T) / 2.0 | |
| try: | |
| return self._rng.multivariate_normal(mean, covariance) | |
| except np.linalg.LinAlgError: | |
| stabilized = covariance + np.eye(covariance.shape[0]) * 1.0e-9 | |
| return self._rng.multivariate_normal(mean, stabilized) | |
| def _prepare_context_array( | |
| self, | |
| values: Sequence[float] | np.ndarray, | |
| *, | |
| include_bias: bool, | |
| ) -> np.ndarray: | |
| """Convert raw values into a clean one-dimensional context vector.""" | |
| raw_array = np.asarray(values, dtype=float).reshape(-1) | |
| if include_bias: | |
| raw_array = np.concatenate([np.asarray([1.0], dtype=float), raw_array]) | |
| if raw_array.size == 0: | |
| raw_array = np.asarray([0.0], dtype=float) | |
| array = np.nan_to_num( | |
| raw_array, | |
| nan=float(self.default_context_value), | |
| posinf=float(self.default_context_value), | |
| neginf=float(self.default_context_value), | |
| ) | |
| if self.clip_context_abs_value is not None: | |
| limit = abs(float(self.clip_context_abs_value)) | |
| array = np.clip(array, -limit, limit) | |
| norm = float(np.linalg.norm(array)) | |
| if self.max_context_norm is not None and norm > float(self.max_context_norm) > 0: | |
| array = array * (float(self.max_context_norm) / max(norm, _EPSILON)) | |
| norm = float(np.linalg.norm(array)) | |
| if self.normalize_context and norm > _EPSILON: | |
| array = array / norm | |
| return array.astype(float) | |
| def _context_from_external( | |
| self, | |
| context: Sequence[float] | Mapping[str, Any] | np.ndarray | None, | |
| ) -> np.ndarray | None: | |
| """Convert an externally supplied context object into a vector.""" | |
| if context is None: | |
| return None | |
| if isinstance(context, Mapping): | |
| raw_vector = context.get("context_vector", context.get("features", _MISSING)) | |
| if raw_vector is not _MISSING and not isinstance(raw_vector, Mapping): | |
| try: | |
| return self._prepare_context_array( | |
| cast(Sequence[float], raw_vector), | |
| include_bias=self.include_bias_for_external_context, | |
| ) | |
| except (TypeError, ValueError): | |
| return None | |
| values: list[float] = [] | |
| if self.include_bias_for_external_context: | |
| values.append(1.0) | |
| keys = self.context_keys or tuple(sorted(str(key) for key in context.keys())) | |
| for key in keys: | |
| raw_value = _get_mapping_path(context, key, context.get(key, _MISSING)) | |
| values.append(_value_as_feature(raw_value, self.default_context_value)) | |
| return self._prepare_context_array(values, include_bias=False) | |
| try: | |
| return self._prepare_context_array( | |
| cast(Sequence[float], context), | |
| include_bias=self.include_bias_for_external_context, | |
| ) | |
| except (TypeError, ValueError): | |
| return None | |
| def _context_for_update( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| behavior: Behavior | str | None, | |
| context: Sequence[float] | Mapping[str, Any] | np.ndarray | None, | |
| arm_key: str, | |
| ) -> np.ndarray | None: | |
| """Resolve the context vector for a reward update.""" | |
| external = self._context_from_external(context) | |
| if external is not None: | |
| return external | |
| if arm_key == self._last_selected_arm_key and self._last_selected_context: | |
| return np.asarray(self._last_selected_context, dtype=float) | |
| if isinstance(behavior, Behavior): | |
| return self.context_vector(agent, world, behavior=behavior) | |
| return None | |
| def _arm_key_for_update(self, behavior: Behavior | str | None) -> str | None: | |
| """Resolve which arm should receive a reward update.""" | |
| if behavior is None: | |
| return self._last_selected_arm_key | |
| if isinstance(behavior, str): | |
| return behavior | |
| return self.arm_key_for_behavior(behavior, behavior_name=self.behavior_name(behavior)) | |
| def _clip_reward(self, reward: float) -> float: | |
| """Clip reward to configured bounds.""" | |
| clipped = float(reward) | |
| if self.reward_clip_min is not None: | |
| clipped = max(float(self.reward_clip_min), clipped) | |
| if self.reward_clip_max is not None: | |
| clipped = min(float(self.reward_clip_max), clipped) | |
| return clipped | |
| def _record_decision( | |
| self, | |
| agent: Agent, | |
| world: World, | |
| selected: BanditEvaluation | None, | |
| evaluations: Sequence[BanditEvaluation], | |
| reason: str, | |
| ) -> None: | |
| """Record a decision trace locally and optionally in agent memory.""" | |
| record = { | |
| "policy_name": self.policy_name, | |
| "agent_id": _agent_id(agent), | |
| "world_step": _world_step(world), | |
| "decision_index": self._decision_count, | |
| "algorithm": _normalize_algorithm(self.algorithm).value, | |
| "reason": reason, | |
| "selected_arm_key": None if selected is None else selected.arm_key, | |
| "selected_behavior_name": None if selected is None else selected.behavior_name, | |
| "selected_score": None if selected is None else selected.score, | |
| "evaluations": [evaluation.to_dict() for evaluation in evaluations], | |
| } | |
| _append_bounded(self._decision_history, copy.deepcopy(record), self.max_decision_history) | |
| if self.write_decisions_to_agent_memory: | |
| memory = _mutable_agent_memory(agent) | |
| history = memory.get(self.decision_memory_key) | |
| if not isinstance(history, MutableSequence): | |
| history = [] | |
| memory[self.decision_memory_key] = history | |
| _append_bounded(history, copy.deepcopy(record), self.max_decision_history) | |
| memory[self.latest_decision_memory_key] = copy.deepcopy(record) | |
| self._decision_count += 1 | |
| Policy = ContextualBanditPolicy | |
| POLICY_REGISTRY: Mapping[str, type[BasePolicy]] = MappingProxyType( | |
| { | |
| ContextualBanditPolicy.name: ContextualBanditPolicy, | |
| } | |
| ) | |
| __all__ = [ | |
| "ArmKeyMode", | |
| "BanditAlgorithm", | |
| "BanditArmState", | |
| "BanditEvaluation", | |
| "ContextualBanditPolicy", | |
| "POLICY_REGISTRY", | |
| "Policy", | |
| ] |