""" 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" @dataclass(frozen=True) 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)), } @dataclass 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)), } @classmethod 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__ @dataclass 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) ] @property def arms(self) -> Mapping[str, BanditArmState]: """Return a read-only view of learned arm states.""" return MappingProxyType(self._arms) @property def last_evaluations(self) -> tuple[BanditEvaluation, ...]: """Return the latest selection evaluations.""" return self._last_evaluations @property def decision_history(self) -> tuple[Mapping[str, Any], ...]: """Return immutable recent decision history.""" return tuple(copy.deepcopy(self._decision_history)) @property 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", ]