WorldSmithAI / policies /contextual_bandit.py
Srishti280992's picture
Upload 39 files
caad8d0 verified
Raw History Blame
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"
@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",
]