Team Ai
Apppublic

EC256/openenv-data-engineering

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py226 linesDownload Raw Back to root
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