| """Gym env: a single-step CONTEXTUAL BANDIT over personas. |
| |
| reset() samples a persona and returns its FEATURE VECTOR. step(action) builds the brief, grades it |
| against a rule-derived answer key, and returns the reward. terminated is always True — one action |
| = one ad = one reward. This is a contextual bandit, not an MDP: the action does not affect any |
| future state (see the spec, §13). |
| |
| Rendering is OFF by default: the reward is computed from the brief alone, so training needs no |
| pixels at all. Turn it on for the demo/UI. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import logging |
| import os |
|
|
| import gymnasium as gym |
| import numpy as np |
| from gymnasium import spaces |
|
|
| from .brief import NVEC, action_to_brief |
| from .config import DEV_CACHE_DIR, EVAL_CACHE_DIR, PROMOTABLE_MODELS |
| from .models import resolve_image_model |
| from .personas import FEATURE_NVEC, HELD_OUT, POOL, TRAIN |
| from .renderer.backgrounds import generate_background |
| from .renderer.cache import Cache |
| from .renderer.overlay import render_ad |
| from .renderer.provider import default_provider |
| from .rules import truth_for |
| from .verifier.reward import compute_reward |
|
|
| logger = logging.getLogger(__name__) |
|
|
| _SPLITS = {"train": TRAIN, "heldout": HELD_OUT, "all": POOL} |
|
|
|
|
| class AdCreativeEnv(gym.Env): |
| metadata = {"render_modes": []} |
|
|
| def __init__( |
| self, |
| *, |
| split: str = "all", |
| mode: str = "placeholder", |
| render: bool = False, |
| model=None, |
| cache=None, |
| provider=None, |
| replay: bool = True, |
| ): |
| super().__init__() |
| if split not in _SPLITS: |
| raise ValueError(f"unknown split {split!r} (use {sorted(_SPLITS)})") |
| self.personas = _SPLITS[split] |
| self.split = split |
| self.render_enabled = render |
|
|
| self.action_space = spaces.MultiDiscrete(np.array(NVEC)) |
| self.observation_space = spaces.MultiDiscrete(np.array(FEATURE_NVEC)) |
| self.current_persona = self.personas[0] |
|
|
| self.mode = mode |
| self.model = resolve_image_model(model) |
| if mode == "provider": |
| if self.model not in PROMOTABLE_MODELS: |
| raise ValueError( |
| f"image model {self.model!r} is not in config.PROMOTABLE_MODELS — a " |
| "non-promotable model cannot be frozen into the eval cache." |
| ) |
| self._cache = cache or Cache(DEV_CACHE_DIR, EVAL_CACHE_DIR, replay=replay) |
| self._provider = provider or default_provider |
| self._replay = replay |
| elif mode != "placeholder": |
| raise ValueError(f"unknown mode {mode!r} (use 'placeholder' or 'provider')") |
|
|
| logger.info( |
| "env init", |
| extra={"mode": self.mode, "model": self.model, "split": split, "render": render}, |
| ) |
|
|
| def _background_source(self): |
| """A persona-BOUND closure. render_ad calls background_source(scene) with the scene |
| alone, so persona context is captured HERE — that is how it can reach the cache key |
| without threading it through Brief or the renderer signature.""" |
| if self.mode != "provider": |
| return None |
|
|
| def _source(scene: str): |
| |
| |
| return generate_background( |
| scene, |
| cache=self._cache, |
| provider=self._provider, |
| model=self.model, |
| api_key=os.environ.get("OPENROUTER_API_KEY"), |
| replay=self._replay, |
| ) |
|
|
| return _source |
|
|
| def reset(self, *, seed=None, options=None): |
| super().reset(seed=seed) |
| idx = int(self.np_random.integers(len(self.personas))) |
| self.current_persona = self.personas[idx] |
| p = self.current_persona |
| context = ( |
| f"{p.trip_type} traveller, {p.budget_tier} budget, on {p.device} " |
| f"({p.time_of_day}) in {p.city}" |
| + (f", attending {p.event} at {p.venue}" if p.event else "") |
| ) |
| return np.array(p.features, dtype=np.int64), {"persona_id": idx, "context": context} |
|
|
| def step(self, action): |
| persona = self.current_persona |
| brief = action_to_brief(action, persona) |
| truth = truth_for(persona) |
|
|
| |
| reward, breakdown = compute_reward(brief, truth) |
|
|
| info = {"breakdown": breakdown, "brief": brief} |
| if self.render_enabled: |
| info["image"] = render_ad(brief, self._background_source()) |
|
|
| logger.debug( |
| "env step", extra={"reward": float(reward), "quality": breakdown["quality"]} |
| ) |
| return np.array(persona.features, dtype=np.int64), float(reward), True, False, info |
|
|