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,
        )