EC256/openenv-data-engineering
0
1import pandas as pd2import json3import os4from typing import Tuple, Dict, Any5 6from models import State, Observation, Reward7from tasks import TASKS8 9 10class DataEnv:11 def __init__(self):12 self._state = None13 self._task = None14 self._df = None15 self._tables = {}16 17 def reset(self, task_id: str = "easy_data_cleaning") -> Observation:18 if task_id not in TASKS:19 raise ValueError(f"Task {task_id} not found.")20 21 self._task = TASKS[task_id]22 self._tables = self._task.load_tables()23 24 primary_table = self._task.get_primary_table()25 self._df = self._tables[primary_table].copy()26 27 self._state = State(28 task_id=task_id,29 step_count=0,30 max_steps=20,31 is_done=False,32 cumulative_reward=0.0,33 tables_loaded=[primary_table],34 )35 36 return self._get_obs(feedback="Environment initialized.")37 38 def _get_obs(self, feedback: str) -> Observation:39 self._df.columns = self._df.columns.astype(str)40 41 safe_df = self._df.head(3).where(pd.notnull(self._df.head(3)), None)42 hints = self._task.get_debug_hints(self._df)43 return Observation(44 dataset_sample=safe_df.to_dict(orient="records"),45 columns=list(self._df.columns),46 dtypes={col: str(dtype) for col, dtype in self._df.dtypes.items()},47 total_rows=len(self._df),48 missing_values=self._df.isna().sum().to_dict(),49 task_description=self._task.description,50 feedback=feedback,51 debug_hints=hints,52 )53 54 def state(self) -> State:55 return self._state56 57 def step(self, action) -> Tuple[Observation, Reward, bool, Dict[str, Any]]:58 self._state.step_count += 159 feedback = ""60 reward_value = 0.061 62 if self._state.is_done:63 return (64 self._get_obs("Already done."),65 Reward(value=0.0, message="Already done."),66 True,67 {},68 )69 70 try:71 action_type = getattr(action, "action_type", None) or action.get(72 "action_type"73 )74 75 if action_type == "submit":76 score, feedback = self._task.grade(self._df, self._tables)77 reward_value = score78 self._state.is_done = True79 80 elif action_type == "drop_column":81 col = getattr(action, "column_name", None)82 if col in self._df.columns:83 self._df.drop(columns=[col], inplace=True)84 feedback = f"Dropped column '{col}'."85 reward_value = 0.0386 else:87 feedback = f"Column '{col}' not found."88 reward_value = -0.0589 90 elif action_type == "rename_column":91 old_n = getattr(action, "old_name", None)92 new_n = getattr(action, "new_name", None)93 if old_n in self._df.columns:94 self._df.rename(columns={old_n: new_n}, inplace=True)95 feedback = f"Renamed '{old_n}' to '{new_n}'."96 reward_value = 0.0297 else:98 feedback = f"Column '{old_n}' not found."99 reward_value = -0.05100 101 elif action_type == "fill_nan":102 col = getattr(action, "column_name", None)103 val = getattr(action, "fill_value", None)104 if col in self._df.columns:105 before_na = int(self._df[col].isna().sum())106 self._df[col] = self._df[col].fillna(val)107 after_na = int(self._df[col].isna().sum())108 feedback = f"Filled NaNs in '{col}' with {val}. Reduced missing by {before_na - after_na}."109 reward_value = 0.03 if after_na < before_na else 0.0110 else:111 feedback = f"Column '{col}' not found."112 reward_value = -0.05113 114 elif action_type == "drop_duplicates":115 subs = getattr(action, "subset", None)116 before = len(self._df)117 self._df.drop_duplicates(subset=subs, inplace=True)118 after = len(self._df)119 dropped = before - after120 feedback = f"Dropped {dropped} duplicate rows."121 reward_value = 0.03 if dropped > 0 else -0.02122 123 elif action_type == "parse_json":124 col = getattr(action, "column_name", None)125 if col in self._df.columns:126 127 def safe_parse(val):128 try:129 if isinstance(val, dict):130 return val131 return json.loads(val)132 except Exception:133 return {}134 135 parsed = self._df[col].apply(safe_parse).apply(pd.Series)136 new_cols = list(parsed.columns)137 # Drop the original column first to avoid duplicates138 self._df = self._df.drop(columns=[col])139 self._df = pd.concat([self._df, parsed], axis=1)140 feedback = f"Parsed JSON from '{col}' into columns: {new_cols}."141 reward_value = 0.08142 else:143 feedback = f"Column '{col}' not found."144 reward_value = -0.05145 146 elif action_type == "cast_type":147 col = getattr(action, "column_name", None)148 target = getattr(action, "target_type", None)149 if col in self._df.columns:150 type_map = {151 "string": str,152 "integer": lambda x: pd.to_numeric(x, errors="coerce").astype(153 "Int64"154 ),155 "float": lambda x: pd.to_numeric(x, errors="coerce"),156 "boolean": lambda x: x.astype(bool),157 "datetime": lambda x: pd.to_datetime(x, errors="coerce"),158 }159 if target in type_map:160 self._df[col] = type_map[target](self._df[col])161 feedback = f"Cast '{col}' to {target}."162 reward_value = 0.05163 else:164 feedback = f"Unknown target type: {target}"165 reward_value = -0.05166 else:167 feedback = f"Column '{col}' not found."168 reward_value = -0.05169 170 elif action_type == "merge_tables":171 left_on = getattr(action, "left_on", None)172 right_on = getattr(action, "right_on", None)173 how = getattr(action, "how", "inner")174 if left_on in self._df.columns:175 merged = False176 for tname, tdf in self._tables.items():177 if (178 right_on in tdf.columns179 and tname not in self._state.tables_loaded180 ):181 self._df = self._df.merge(182 tdf, left_on=left_on, right_on=right_on, how=how183 )184 self._state.tables_loaded.append(tname)185 feedback = f"Merged with '{tname}' on {left_on}={right_on} (how={how}). New shape: {len(self._df)} rows."186 reward_value = 0.1187 merged = True188 break189 if not merged:190 feedback = (191 f"No available table has column '{right_on}' to merge on."192 )193 reward_value = -0.05194 else:195 feedback = f"Column '{left_on}' not found in current dataframe."196 reward_value = -0.05197 198 elif action_type == "execute_pandas":199 code = getattr(action, "code", None)200 local_vars = {"df": self._df, "pd": pd, "json": json}201 exec(code, {}, local_vars)202 self._df = local_vars["df"]203 feedback = "Successfully executed Pandas code."204 reward_value = 0.02205 206 else:207 feedback = f"Unknown action type: {action_type}"208 reward_value = -0.1209 210 except Exception as e:211 feedback = f"Error executing action: {str(e)}"212 reward_value = -0.1213 214 self._state.cumulative_reward += reward_value215 216 if self._state.step_count >= self._state.max_steps:217 self._state.is_done = True218 final_score, final_grade_feedback = self._task.grade(self._df, self._tables)219 reward_value = final_score220 feedback += f"\nMax steps reached. Final Grade: {final_grade_feedback}"221 222 obs = self._get_obs(feedback)223 reward = Reward(value=max(-1.0, min(1.0, reward_value)), message=feedback)224 225 return obs, reward, self._state.is_done, {}226 