ParthKulshreshtha's picture
Deploy personalized ad environment
526cf2e verified
Raw
History Blame Contribute Delete
4.86 kB
"""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