"""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): # NOTE: no `city=` — generate_background does not accept one (Task 5 deferred; the # 4 promoted non-city backgrounds are what the eval cache holds). 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) # The reward reads the BRIEF ONLY. No pixels. 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