AdamK29/Meta-OpenENV-Hackathon
0
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