Spaces:
Sleeping
Sleeping
File size: 22,698 Bytes
41ea373 98b952a ebae6ab 41ea373 258783b 41ea373 0e4f105 41ea373 5c28dc0 41ea373 258783b 5c28dc0 41ea373 0e4f105 dfa9070 998d987 09f7d63 79cb04a 41ea373 258783b ebae6ab 41ea373 ebae6ab 41ea373 ebae6ab 41ea373 ebae6ab 98b952a 41ea373 258783b 5c28dc0 41ea373 0e4f105 dfa9070 998d987 09f7d63 79cb04a 41ea373 dfa9070 998d987 09f7d63 41ea373 ebae6ab 98b952a ebae6ab 41ea373 ebae6ab 41ea373 ebae6ab 41ea373 ebae6ab 09f7d63 998d987 258783b 5c28dc0 41ea373 258783b 5c28dc0 41ea373 ebae6ab 41ea373 dfa9070 41ea373 dfa9070 0e4f105 41ea373 98b952a 41ea373 98b952a 41ea373 ebae6ab 98b952a 258783b 0e4f105 98b952a 41ea373 998d987 258783b 5c28dc0 dfa9070 998d987 79cb04a 41ea373 258783b 5c28dc0 dfa9070 998d987 79cb04a 0e4f105 41ea373 258783b 41ea373 258783b 41ea373 258783b 41ea373 5c28dc0 0e4f105 41ea373 258783b 41ea373 ebae6ab 09f7d63 98b952a 41ea373 258783b 5c28dc0 0e4f105 dfa9070 98b952a 41ea373 09f7d63 ebae6ab 41ea373 258783b dfa9070 98b952a 41ea373 5c28dc0 09f7d63 41ea373 5c28dc0 dfa9070 09f7d63 41ea373 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 | import json
import random
import time
from collections import Counter
from typing import Optional, Tuple
from viral_script_engine.agents.critic import CriticAgent
from viral_script_engine.agents.defender import DefenderAgent
from viral_script_engine.agents.rewriter import RewriterAgent
from viral_script_engine.agents.reasoning_parser import ReasoningParser, ArbitratorParseError
from viral_script_engine.environment.actions import ArbitratorAction
from viral_script_engine.environment.episode_state import EpisodeState
from viral_script_engine.environment.observations import (
DebateRound, Observation, RewardComponents,
)
from viral_script_engine.agents.moderation_agent import ModerationAgent
from viral_script_engine.agents.originality_agent import OriginalityAgent
from viral_script_engine.rewards.r1_hook_strength import HookStrengthReward
from viral_script_engine.rewards.r2_coherence import CoherenceReward
from viral_script_engine.rewards.r3_cultural_alignment import CulturalAlignmentReward
from viral_script_engine.rewards.r4_debate_resolution import DebateResolutionReward
from viral_script_engine.rewards.r5_defender_preservation import DefenderPreservationReward
from viral_script_engine.rewards.r6_safety import SafetyReward
from viral_script_engine.rewards.r7_originality import OriginalityReward
from viral_script_engine.rewards.reward_aggregator import RewardAggregator
from viral_script_engine.rewards.process_reward import ProcessReward, ProcessRewardResult
from viral_script_engine.personas.creator_profile import CreatorProfile, CreatorTier
from viral_script_engine.personas.profile_generator import ProfileGenerator
from viral_script_engine.rewards.r8_persona_fit import PersonaFitReward
from viral_script_engine.rewards.r9_platform_pacing import PlatformPacingReward
from viral_script_engine.platforms.platform_spec import PlatformRegistry
from viral_script_engine.memory.memory_compressor import MemoryCompressor
from viral_script_engine.memory.history_store import HistoryStore
from viral_script_engine.rewards.r10_retention_curve import RetentionCurveReward
_TIERS = {
"easy": ["S01", "S02", "S03", "S04"],
"medium": ["S05", "S06", "S07"],
"hard": ["S08", "S09", "S10"],
"self_generated": [],
}
class ViralScriptEnv:
def __init__(
self,
scripts_path: str = "data/test_scripts/scripts.json",
max_steps: int = 5,
difficulty: str = "easy",
use_anti_gaming: bool = True,
cultural_kb_path: str = "data/cultural_kb.json",
use_escalation: bool = True,
difficulty_tracker=None,
escalation_engine=None,
):
self.max_steps = max_steps
self.difficulty = difficulty
self.use_anti_gaming = use_anti_gaming
self.use_escalation = use_escalation
with open(scripts_path) as f:
all_scripts = json.load(f)
tier_ids = _TIERS.get(difficulty, [])
self._scripts = [s for s in all_scripts if s["script_id"] in tier_ids]
if not self._scripts:
self._scripts = all_scripts
self.critic = CriticAgent(backend="hf")
self.defender = DefenderAgent(backend="hf")
self.rewriter = RewriterAgent(backend="hf")
self.r1 = HookStrengthReward()
self.r2 = CoherenceReward()
self.r3 = CulturalAlignmentReward(knowledge_base_path=cultural_kb_path)
self.r4 = DebateResolutionReward(critic_agent=self.critic)
self.r5 = DefenderPreservationReward()
self.r6 = SafetyReward()
self.r7 = OriginalityReward()
self.moderation_agent = ModerationAgent()
self.originality_agent = OriginalityAgent()
self.aggregator = RewardAggregator()
self.reasoning_parser = ReasoningParser()
self.process_reward_calc = ProcessReward()
self.profile_generator = ProfileGenerator()
self.r8 = PersonaFitReward()
self.r9 = PlatformPacingReward()
self.platform_registry = PlatformRegistry()
self.memory_compressor = MemoryCompressor()
self.history_store = HistoryStore()
self.r10 = RetentionCurveReward(cultural_kb_path=cultural_kb_path)
self._state: Optional[EpisodeState] = None
self._current_profile: Optional[CreatorProfile] = None
self._current_platform: str = "Reels"
self._current_creator_id: str = "default"
self._current_history_buffer = None
if use_escalation:
if difficulty_tracker is None:
from viral_script_engine.escalation.difficulty_tracker import DifficultyTracker
difficulty_tracker = DifficultyTracker()
if escalation_engine is None:
from viral_script_engine.escalation.critic_escalation_engine import CriticEscalationEngine
escalation_engine = CriticEscalationEngine()
self.difficulty_tracker = difficulty_tracker
self.escalation_engine = escalation_engine
# Track first-step critic output per episode for dominant class detection
self._first_critique = None
self._timeout_count: int = 0
def reset_from_config(self, episode_config: dict) -> Tuple[dict, dict]:
"""Reset the environment to a specific episode config from curriculum JSONL."""
script = {
"script_id": episode_config.get("script_id", "unknown"),
"script_text": episode_config["script_text"],
"region": episode_config["region"],
"platform": episode_config["platform"],
"niche": episode_config["niche"],
}
return self._reset_with_script(script, episode_config.get("difficulty", self.difficulty))
def reset(self, seed=None, options=None) -> Tuple[dict, dict]:
if seed is not None:
random.seed(seed)
self._first_critique = None
used_escalation = False
if self.use_escalation and self.difficulty_tracker and self.escalation_engine:
mastered = self.difficulty_tracker.get_mastered_classes()
if mastered:
challenge = self.escalation_engine.get_next_challenge(self.difficulty_tracker)
if challenge is None:
# Generate a new escalated challenge from the first mastered class
src_class = mastered[0]
example_script = random.choice(self._scripts)
challenge = self.escalation_engine.escalate(
mastered_class=src_class,
original_script_example=example_script["script_text"],
region=example_script.get("region", "pan_india_english"),
platform=example_script.get("platform", "Reels"),
)
script = challenge.to_script_dict()
print(f"[ESCALATION] Using self-generated challenge for class '{challenge.source_class}' — {challenge.why_its_harder}")
obs, info = self._reset_with_script(script, "self_generated")
info["escalation_used"] = True
info["escalation_source_class"] = challenge.source_class
return obs, info
script = random.choice(self._scripts)
obs, info = self._reset_with_script(script, self.difficulty)
info["escalation_used"] = False
return obs, info
def _reset_with_script(self, script: dict, difficulty: str) -> Tuple[dict, dict]:
self._current_creator_id = script.get("creator_id", script.get("script_id", "default"))
self._current_history_buffer = self.history_store.load(self._current_creator_id)
self._current_platform = script.get("platform", "Reels")
r1_result = self.r1.score(script["script_text"], platform=self._current_platform)
r2_result = self.r2.score(script["script_text"], script["script_text"], platform=self._current_platform)
r3_result = self.r3.score(script["script_text"], script.get("region", "pan_india_english"))
mod_out = self.moderation_agent.check(script["script_text"])
orig_out = self.originality_agent.check(script["script_text"])
r6_result = self.r6.score(mod_out)
r7_result = self.r7.score(orig_out)
initial_rewards = RewardComponents(
r1_hook_strength=r1_result.score,
r2_coherence=r2_result.score,
r3_cultural_alignment=r3_result.score,
r6_safety=r6_result.score,
r7_originality=r7_result.score,
)
initial_rewards.compute_total()
self._state = EpisodeState.new(
script=script,
max_steps=self.max_steps,
difficulty_level=difficulty,
initial_rewards=initial_rewards,
)
# Phase 8: generate a creator profile matching episode difficulty
self._current_profile = self._generate_profile_for_difficulty(
difficulty=difficulty,
niche=script.get("niche", "personal finance"),
seed=hash(script.get("script_id", "default")) % (2 ** 31),
)
return self._build_observation().model_dump(), {}
def _generate_profile_for_difficulty(
self, difficulty: str, niche: str, seed: int
) -> CreatorProfile:
"""Map episode difficulty to an appropriate creator tier."""
tier_map = {
"easy": [CreatorTier.BEGINNER, CreatorTier.GROWING],
"medium": [CreatorTier.GROWING, CreatorTier.ESTABLISHED],
"hard": [CreatorTier.ESTABLISHED, CreatorTier.VERIFIED],
"self_generated": [CreatorTier.ESTABLISHED, CreatorTier.VERIFIED],
}
import random as _rng
tiers = tier_map.get(difficulty, [CreatorTier.GROWING])
tier = _rng.Random(seed).choice(tiers)
return self.profile_generator.generate(tier=tier, niche=niche, seed=seed)
def step(self, action: dict, raw_output: str = None) -> Tuple[dict, float, bool, bool, dict]:
if self._state is None:
raise RuntimeError("Call reset() before step()")
_step_start = time.time()
arb_action = ArbitratorAction(**action)
try:
critique = self.critic.critique(
script=self._state.current_script,
region=self._state.region,
platform=self._state.platform,
niche=self._state.niche,
)
except TimeoutError:
self._timeout_count += 1
info = {"timeout": True, "timeout_agent": "critic", "timeout_count": self._timeout_count}
return self._build_observation().model_dump(), 0.0, False, True, info
# Track first critique for dominant class detection at episode end
if self._state.step_num == 0:
self._first_critique = critique
try:
defender_output = self.defender.defend(
script=self._state.current_script,
critic_claims=critique.claims,
region=self._state.region,
platform=self._state.platform,
)
except TimeoutError:
self._timeout_count += 1
info = {"timeout": True, "timeout_agent": "defender", "timeout_count": self._timeout_count}
return self._build_observation().model_dump(), 0.0, False, True, info
# Phase 7: parse reasoning chain and compute process reward before rewrite
reasoning_chain = None
process_result = None
if raw_output:
try:
reasoning_chain = self.reasoning_parser.parse(raw_output)
process_result = self.process_reward_calc.score(
reasoning_chain=reasoning_chain,
critic_claims=critique.claims,
defender_output=defender_output,
current_reward_components=self._state.last_reward_components,
episode_start_components=self._state.episode_start_rewards,
)
except ArbitratorParseError:
reasoning_chain = None
process_result = None
try:
rewrite_result = self.rewriter.rewrite(self._state.current_script, arb_action)
except TimeoutError:
self._timeout_count += 1
info = {"timeout": True, "timeout_agent": "rewriter", "timeout_count": self._timeout_count}
return self._build_observation().model_dump(), 0.0, False, True, info
new_script = rewrite_result.rewritten_script
r1_result = self.r1.score(new_script, platform=self._current_platform)
r2_result = self.r2.score(self._state.original_script, new_script, platform=self._current_platform)
r3_result = self.r3.score(new_script, self._state.region)
targeted_claim = next(
(c for c in critique.claims if c.claim_id == arb_action.critique_claim_id),
critique.claims[0] if critique.claims else None,
)
r4_result = self.r4.score(
new_script=new_script,
original_action=arb_action,
original_claim=targeted_claim,
region=self._state.region,
platform=self._state.platform,
niche=self._state.niche,
) if targeted_claim else None
r5_result = self.r5.score(defender_output, new_script)
moderation_out = self.moderation_agent.check(new_script)
originality_out = self.originality_agent.check(new_script)
r6_result = self.r6.score(moderation_out)
r7_result = self.r7.score(originality_out)
# Phase 8: compute R8 persona fit
r8_score = None
if self._current_profile is not None and targeted_claim is not None:
r8_result = self.r8.score(
action=arb_action,
creator_profile=self._current_profile,
addressed_critique_class=targeted_claim.critique_class,
)
r8_score = r8_result.score
# Phase 9: compute R9 platform pacing
r9_result = self.r9.score(new_script, platform=self._current_platform)
# Phase 12: compute R10 retention curve reward
r10_score = None
if self.r10.predictor._trained:
try:
r10_result = self.r10.score(
original_script=self._state.original_script,
rewritten_script=new_script,
platform=self._current_platform,
region=self._state.region,
action_type=str(arb_action.action_type.value),
episode_id=self._state.episode_id,
)
r10_score = r10_result.score
except Exception:
r10_score = None
components = RewardComponents(
r1_hook_strength=r1_result.score,
r2_coherence=r2_result.score,
r3_cultural_alignment=r3_result.score,
r4_debate_resolution=r4_result.score if r4_result else None,
r5_defender_preservation=r5_result.score,
r6_safety=r6_result.score,
r7_originality=r7_result.score,
r8_persona_fit=r8_score,
r9_platform_pacing=r9_result.score,
r10_retention_curve=r10_score,
process_reward=process_result.weighted_contribution if process_result else None,
)
self._state.action_history.append(arb_action.action_type)
if self.use_anti_gaming:
components, anti_log = self.aggregator.compute(
components,
self._state.episode_start_rewards,
self._state.action_history,
episode_id=self._state.episode_id,
step_num=self._state.step_num,
)
else:
components.compute_total()
from viral_script_engine.rewards.reward_aggregator import AntiGamingLog
anti_log = AntiGamingLog(
episode_id=self._state.episode_id,
step_num=self._state.step_num,
triggered=False,
penalty_applied=0.0,
pre_penalty_total=components.total,
post_penalty_total=components.total,
)
round_ = DebateRound(
step_num=self._state.step_num,
critic_claims=critique.claims,
defender_response=defender_output.model_dump(),
arbitrator_action=arb_action,
rewrite_diff=rewrite_result.diff,
reward_components=components,
moderation_output=moderation_out.model_dump(),
originality_output=originality_out.model_dump(),
reasoning_chain=reasoning_chain.model_dump() if reasoning_chain else None,
)
self._state.debate_history.append(round_)
self._state.current_script = new_script
self._state.last_reward_components = components
self._state.step_num += 1
if not hasattr(self._state, "anti_gaming_logs"):
self._state.anti_gaming_logs = []
self._state.anti_gaming_logs.append(anti_log.model_dump())
terminated = (
self._state.step_num >= self._state.max_steps
or components.total >= 0.9
)
if terminated and self.use_escalation and self.difficulty_tracker:
dominant_class = self._get_dominant_critique_class()
r4_score = components.r4_debate_resolution if components.r4_debate_resolution is not None else 0.0
self.difficulty_tracker.record_episode(
dominant_critique_class=dominant_class,
r4_score=r4_score,
episode_id=self._state.episode_id,
)
if terminated:
episode_number = (
(self._current_history_buffer.total_episodes + 1)
if self._current_history_buffer else 1
)
new_memory = self.memory_compressor.compress(
episode_log=self._build_episode_log(),
episode_number=episode_number,
)
self._current_history_buffer = self.memory_compressor.update_buffer(
self._current_history_buffer, new_memory, self._current_creator_id
)
self.history_store.save(self._current_history_buffer)
if time.time() - _step_start > 120:
self._timeout_count += 1
return self._build_observation().model_dump(), 0.0, False, True, {
"timeout": True, "timeout_agent": "step_wall_clock",
"timeout_count": self._timeout_count,
}
info = {
"reward_components": components.model_dump(),
"anti_gaming_triggered": anti_log.triggered,
"penalty_reason": anti_log.rule_triggered,
"anti_gaming_log": anti_log.model_dump(),
"moderation_output": moderation_out.model_dump(),
"originality_output": originality_out.model_dump(),
"process_reward_result": process_result.model_dump() if process_result else None,
"reasoning_chain": reasoning_chain.model_dump() if reasoning_chain else None,
"creator_profile": self._current_profile.model_dump(mode="json") if self._current_profile else None,
"timeout_count": self._timeout_count,
}
return self._build_observation().model_dump(), components.total, terminated, False, info
def _build_episode_log(self) -> dict:
s = self._state
first_claims = []
if self._first_critique and self._first_critique.claims:
first_claims = [c.model_dump() for c in self._first_critique.claims]
return {
"episode_id": s.episode_id,
"niche": s.niche,
"platform": s.platform,
"actions_taken": [a.value if hasattr(a, "value") else str(a) for a in s.action_history],
"first_critique_claims": first_claims,
"initial_reward_components": s.episode_start_rewards.model_dump(),
"final_reward_components": s.last_reward_components.model_dump(),
"final_total_reward": s.last_reward_components.total,
}
def _get_dominant_critique_class(self) -> str:
"""Return the most common critique_class from the first episode critique."""
if self._first_critique is None or not self._first_critique.claims:
return "hook_weakness"
counts = Counter(c.critique_class for c in self._first_critique.claims)
return counts.most_common(1)[0][0]
def state(self) -> dict:
if self._state is None:
return {}
s = self._state
return {
"current_script": s.current_script,
"original_script": s.original_script,
"debate_history": [r.model_dump() for r in s.debate_history],
"reward_components": s.last_reward_components.model_dump(),
"step_num": s.step_num,
"difficulty_level": s.difficulty_level,
"episode_id": s.episode_id,
"anti_gaming_logs": getattr(s, "anti_gaming_logs", []),
"creator_profile": self._current_profile.model_dump(mode="json") if self._current_profile else None,
"timeout_count": self._timeout_count,
}
def _build_observation(self) -> Observation:
s = self._state
last_round = s.debate_history[-1] if s.debate_history else None
mod_flags = []
orig_flags = []
if last_round and last_round.moderation_output:
mod_flags = last_round.moderation_output.get("flags", [])
if last_round and last_round.originality_output:
orig_flags = last_round.originality_output.get("flags", [])
history_context = (
self._current_history_buffer.to_prompt_context()
if self._current_history_buffer else None
)
return Observation(
current_script=s.current_script,
original_script=s.original_script,
region=s.region,
platform=s.platform,
niche=s.niche,
step_num=s.step_num,
max_steps=s.max_steps,
debate_history=s.debate_history,
reward_components=s.last_reward_components,
difficulty_level=s.difficulty_level,
episode_id=s.episode_id,
current_moderation_flags=mod_flags,
current_originality_flags=orig_flags,
creator_profile=self._current_profile.model_dump(mode="json") if self._current_profile else None,
creator_history=self._current_history_buffer.model_dump() if self._current_history_buffer else None,
history_context=history_context,
)
|