CoolFace
Apppublic

BHaritha/msme-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
environment.py278 linesDownload Raw Back to env
1"""2environment.py - MSME Payment Dispute OpenEnv Environment3Implements reset(), step(), state() per OpenEnv spec.4 5FIX 1: Sequential episode mode — tasks chain together.6        Wrong Task 1 classification bleeds into Task 3 context,7        making this a true multi-step environment.8"""9from .tasks import get_scenario10from .graders import grade_task1, grade_task2, grade_task3, grade_task3_multiturn11 12class MSMEDisputeEnv:13    def __init__(self):14        self._mode       = "single"   # "single" or "sequential"15        self._task_id    = None16        self._scenario   = None17        self._done       = False18        self._step_count = 019        self._seed       = None20        self._last_reward     = None21        self._episode_rewards = []22 23        # Sequential mode state24        self._seq_step        = 0     # which task we're on (1,2,3)25        self._seq_scenarios   = {}    # {1: scenario, 2: scenario, 3: scenario}26        self._seq_results     = {}    # {1: grader_result, 2: grader_result}27        self._agent_t1_label  = None  # agent's Task1 answer bleeds into Task328        self._agent_t2_facts  = None  # agent's Task2 answer bleeds into Task329        self._t1_fallback_used = False # True if agent gave invalid Task1 label30 31        # Multi-turn Task 3 state (up to 3 turns: draft → revise → final)32        self._t3_turn          = 133        self._t3_scenario      = None34        self._t3_done          = False35 36    # ─────────────────────────────────────────37    # PUBLIC API38    # ─────────────────────────────────────────39 40    def reset(self, task_id: int = 1, seed: int = None, mode: str = "single") -> dict:41        """42        Start a fresh episode.43 44        mode="single"     — one task, one step (original behaviour)45        mode="sequential" — all 3 tasks chained, agent's earlier46                            answers affect later task context47        task_id           — used in single mode only48        seed              — fixed seed for reproducibility49        """50        self._seed            = seed51        self._done            = False52        self._step_count      = 053        self._last_reward     = None54        self._episode_rewards = []55        self._agent_t1_label  = None56        self._agent_t2_facts  = None57        self._seq_results     = {}58        self._t1_fallback_used = False59        self._t3_turn          = 160        self._t3_done          = False61 62        if mode == "sequential":63            self._mode      = "sequential"64            self._task_id   = None65            self._seq_step  = 166            # Pre-load all 3 scenarios with same seed for consistency67            self._seq_scenarios = {68                1: get_scenario(1, seed),69                2: get_scenario(2, seed),70                3: get_scenario(3, seed),71            }72            self._scenario = self._seq_scenarios[1]73            return self._build_obs(task_id=1)74 75        else:76            if task_id not in (1, 2, 3):77                raise ValueError("task_id must be 1, 2, or 3")78            self._mode    = "single"79            self._task_id = task_id80            self._scenario = get_scenario(task_id, seed)81            return self._build_obs(task_id=task_id)82 83    def step(self, action: dict) -> dict:84        """85        Submit an action.86 87        Single mode:     one step → done.88        Sequential mode: three steps → done after Task 3.89                         Agent's Task1 label and Task2 facts are90                         injected into the Task3 observation so91                         a wrong classification genuinely hurts92                         the final letter quality score.93        """94        if self._scenario is None:95            raise RuntimeError("Call reset() before step()")96        if self._done:97            raise RuntimeError("Episode done. Call reset() to start a new one.")98 99        if self._mode == "sequential":100            return self._step_sequential(action)101        else:102            return self._step_single(action)103 104    def state(self) -> dict:105        """Return current environment state."""106        if self._scenario is None:107            return {"status": "not_started"}108 109        base = {110            "mode":         self._mode,111            "seed":         self._seed,112            "step_count":   self._step_count,113            "done":         self._done,114            "last_reward":  self._last_reward,115        }116 117        if self._mode == "sequential":118            base.update({119                "current_task":    self._seq_step if not self._done else "done",120                "episode_rewards": self._seq_results,121                "average_reward":  (122                    round(sum(r["score"] for r in self._seq_results.values()) /123                          len(self._seq_results), 3)124                    if self._seq_results else None125                ),126                "observation": self._build_obs(self._seq_step) if not self._done else None,127            })128        else:129            base.update({130                "task_id":      self._task_id,131                "scenario_id":  self._scenario.get("id"),132                "observation":  self._build_obs(self._task_id),133            })134 135        return base136 137    # ─────────────────────────────────────────138    # PRIVATE HELPERS139    # ─────────────────────────────────────────140 141    def _step_single(self, action: dict) -> dict:142        result = self._grade(action, self._task_id, self._scenario)143        self._last_reward = result["score"]144        self._step_count += 1145 146        # Multi-turn: Task 3 allows up to 3 turns (draft → revise → final)147        if self._task_id == 3 and result.get("needs_revision") and self._t3_turn < 3:148            self._t3_turn += 1149            done = False  # not done yet — agent should revise150            self._done = False151        else:152            self._t3_done = True153            done = True154            self._done = True155 156        return {157            "state":    self.state(),158            "reward":   result["score"],159            "done":     done,160            "info":     result,161            "feedback": result.get("feedback", []),162            "message":  result.get("message", "")163        }164 165    def _step_sequential(self, action: dict) -> dict:166        current = self._seq_step167        scenario = self._seq_scenarios[current]168 169        result = self._grade(action, current, scenario)170        self._seq_results[current] = result171        self._last_reward = result["score"]172        self._step_count += 1173 174        # Store agent answers so they bleed into Task 3 context.175        # If Task 1 label is invalid/broken, fall back to "delayed_payment"176        # (most common type) so Task 3 context is always gradeable.177        VALID_LABELS = {"delayed_payment", "partial_payment", "payment_denial"}178        if current == 1:179            raw_label = str(action.get("label", "")).strip().lower()180            self._agent_t1_label = raw_label if raw_label in VALID_LABELS else "delayed_payment"181            if raw_label not in VALID_LABELS:182                # Record that a fallback was used — this is visible in state()183                self._t1_fallback_used = True184        elif current == 2:185            self._agent_t2_facts = action186 187        # Advance or finish188        if current < 3:189            self._seq_step += 1190            self._scenario = self._seq_scenarios[self._seq_step]191            done = False192        else:193            self._done = True194            done = True195 196        avg = round(197            sum(r["score"] for r in self._seq_results.values()) / len(self._seq_results), 3198        )199 200        return {201            "state":          self.state(),202            "reward":         result["score"],203            "done":           done,204            "task_completed": current,205            "next_task":      current + 1 if not done else None,206            "episode_avg":    avg if done else None,207            "info":           result,208        }209 210    def _build_obs(self, task_id: int) -> dict:211        if self._mode == "sequential":212            s = self._seq_scenarios.get(task_id, {})213        else:214            s = self._scenario215 216        if task_id == 1:217            return {218                "task":        "classify_dispute",219                "description": "Classify the dispute type from the email.",220                "email":       s["email"],221                "valid_labels": ["delayed_payment", "partial_payment", "payment_denial"],222                "action_format": {"label": "<one of the valid_labels>"},223                "note": "Your classification here will affect the context given in Task 3."224            }225 226        elif task_id == 2:227            return {228                "task":        "extract_facts",229                "description": "Extract structured facts from the formal notice.",230                "email":       s["email"],231                "action_format": {232                    "claimant":     "<company name>",233                    "opponent":     "<company name>",234                    "amount":       "<integer in rupees>",235                    "due_date":     "<date string>",236                    "days_overdue": "<integer>"237                },238                "note": "Your extracted facts will be used as context in Task 3."239            }240 241        elif task_id == 3:242            ctx = dict(s["context"])  # copy243 244            # FIX 1 CORE: Inject agent's earlier answers into Task 3 context.245            # If agent got Task 1 wrong, they get the wrong dispute_type injected.246            # This means a wrong Task 1 genuinely hurts Task 3 letter quality.247            if self._mode == "sequential":248                if self._agent_t1_label:249                    ctx["dispute_type_from_agent"] = self._agent_t1_label250                    ctx["dispute_type"] = self._agent_t1_label  # overwrites ground truth251                if self._agent_t2_facts:252                    for field in ["claimant", "opponent", "amount", "due_date", "days_overdue"]:253                        if field in self._agent_t2_facts:254                            ctx[field] = self._agent_t2_facts[field]255 256            return {257                "task":        "draft_demand_letter",258                "description": "Draft a formal MSME payment demand letter using the context below.",259                "context":     ctx,260                "action_format": {"letter": "<full demand letter text, minimum 150 words>"},261                "note": "Context was built from your Task 1 and Task 2 answers." if self._mode == "sequential" else ""262            }263 264    def _grade(self, action: dict, task_id: int, scenario: dict) -> dict:265        if task_id == 1:266            return grade_task1(action, {"label": scenario["label"]})267        elif task_id == 2:268            return grade_task2(action, scenario["ground_truth"])269        elif task_id == 3:270            # Build grading scenario with possibly agent-modified context271            grading_scenario = dict(scenario)272            if self._mode == "sequential" and self._agent_t1_label:273                ctx = dict(scenario["context"])274                ctx["dispute_type"] = self._agent_t1_label275                grading_scenario = {**scenario, "context": ctx}276            # Use multi-turn grader — tracks which revision turn we're on277            return grade_task3_multiturn(action, grading_scenario, turn=self._t3_turn)278