CoolFace
Apppublic

AdamK29/Meta-OpenENV-Hackathon

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py183 linesDownload Raw Back to server
1import uuid2import random3from openenv.core.env_server import Environment4from core.models import EmailAction, EmailObservation, EmailState5from core.dataset import load_dataset6 7 8class EmailEnv(Environment):9 10    def __init__(self):11        self.state_obj = EmailState()12 13        # RL STATE14        self.inbox = []15        self.current_email = None16        self.done = False17 18        self.total_reward = 0.019        self.missed_urgent = 020        self.processed = 021 22        # dataset cache23        self.dataset = load_dataset()24 25    # ---------------- RESET ----------------26 27    def reset(self, **kwargs):28        task = kwargs.get("task", "easy")29 30        # ---- DATA SPLIT ----31        dataset_size = len(self.dataset)32        33        # Split ratios34        easy_end = int(0.2 * dataset_size)35        medium_end = int(0.6 * dataset_size)36        37        if task == "easy":38            emails = self.dataset[:easy_end]39        40        elif task == "medium":41            emails = self.dataset[easy_end:medium_end]42        43        else:  # hard44            emails = self.dataset[medium_end:]45 46        emails = emails.copy()47        random.shuffle(emails)48 49        self.inbox = emails50        self.current_email = self.inbox.pop(0)51 52        self.done = False53        self.total_reward = 0.054        self.missed_urgent = 055        self.processed = 056 57        self.state_obj = EmailState(58            episode_id=str(uuid.uuid4()),59            step_count=0,60            current_index=0,61            total_emails=len(emails),62            task_type=task,63            score=0.064        )65 66        return self._build_obs(0.0, "Start episode")67 68    # ---------------- STEP ----------------69 70    def step(self, action: EmailAction, **kwargs):71 72        if self.done:73            return self._final_obs(0.01)74    75        self.state_obj.step_count += 176        self.processed += 177    78        gt = self.current_email.get("label")79    80        # -------- NORMALIZED REWARD SYSTEM (STRICT 0 < reward < 1) --------81    82        score = 0.083    84        # 1️⃣ Classification score (0.1 → 0.4)85        if action.content == gt:86            score += 0.487        else:88            score += 0.189    90        # 2️⃣ Urgent handling (0 → 0.3)91        if gt == "urgent":92            if action.content == "urgent":93                score += 0.394            else:95                score += 0.096                self.missed_urgent += 197        else:98            score += 0.2  # non-urgent handled safely99    100        # 3️⃣ Progress reward (0 → 0.2)101        total = self.state_obj.total_emails + 1102        progress = 1.0 - (len(self.inbox) / total)103        score += 0.2 * progress104    105        # 4️⃣ Efficiency (0 → 0.1)106        efficiency = max(0.0, 1.0 - (self.state_obj.step_count / total))107        score += 0.1 * efficiency108    109        # -------- CLAMP FINAL REWARD --------110    111        reward = min(max(score, 0.01), 0.99)112    113        self.total_reward += reward114        self.state_obj.score = self.total_reward115    116        # -------- TRANSITION --------117    118        if len(self.inbox) > 0:119            self.current_email = self.inbox.pop(0)120        else:121            self.done = True122    123        # -------- TERMINATION CONDITIONS --------124    125        # Too many urgent misses → soft penalty126        if self.missed_urgent >= 3:127            self.done = True128            reward = 0.05  # still valid range129    130        # Episode completion reward (normalized)131        if self.done:132            completion = 1.0 - (self.missed_urgent / (self.processed + 1))133            reward = min(max(completion, 0.01), 0.99)134            self.total_reward += reward135    136        # ---------------------------------------137    138        return self._build_obs(reward, "Step complete")139 140    # ---------------- OBS BUILDER ----------------141 142    def _build_obs(self, reward, msg):143 144        urgency_ratio = self._urgent_ratio()145 146        return EmailObservation(147            done=self.done,148            reward=reward,149            email_text=self.current_email["text"] if not self.done else "Inbox Cleared",150            sender="real_user@enron.com",151            subject="Email",152            history=[153                f"processed:{self.processed}",154                f"missed_urgent:{self.missed_urgent}",155                f"total_reward:{self.total_reward:.2f}"156            ],157            message=f"{msg} | inbox:{len(self.inbox)} | urgency_ratio:{urgency_ratio:.2f}"158        )159 160    # ---------------- HELPERS ----------------161 162    def _urgent_ratio(self):163        if len(self.inbox) == 0:164            return 0.0165        urgent = sum(1 for e in self.inbox if e.get("label") == "urgent")166        return urgent / len(self.inbox)167 168    def _final_obs(self, reward):169        return EmailObservation(170            done=True,171            reward=reward,172            email_text="Episode finished",173            sender="system",174            subject="Done",175            history=[],176            message="All emails processed"177        )178 179    # ---------------- STATE ----------------180 181    @property182    def state(self):183        return self.state_obj