File size: 5,541 Bytes
e4f7ea7
 
022e674
 
 
 
 
 
 
 
 
 
 
 
e4f7ea7
022e674
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e4f7ea7
 
022e674
 
e4f7ea7
 
 
 
022e674
 
 
 
 
 
 
 
e4f7ea7
 
 
 
 
 
 
 
 
 
022e674
e4f7ea7
 
 
 
 
 
 
 
 
022e674
e4f7ea7
022e674
 
 
 
 
 
e4f7ea7
 
 
 
 
 
022e674
e4f7ea7
022e674
e4f7ea7
 
 
 
 
 
 
 
 
 
 
 
022e674
e4f7ea7
 
 
 
 
022e674
e4f7ea7
 
 
 
022e674
 
e4f7ea7
 
022e674
 
 
e4f7ea7
 
 
022e674
e4f7ea7
 
022e674
e4f7ea7
022e674
e4f7ea7
022e674
e4f7ea7
 
 
 
 
022e674
e4f7ea7
022e674
 
 
e4f7ea7
 
 
 
 
 
 
 
 
 
 
 
 
 
022e674
 
 
e4f7ea7
 
 
022e674
e4f7ea7
 
022e674
 
e4f7ea7
 
022e674
 
 
 
e4f7ea7
 
 
022e674
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
from env.data_generator import generate_clean_dataset
from env.issue_injector import inject_issues
from env.actions import (
    drop_correlated_feature,
    fill_nulls,
    remove_nulls,
    convert_types,
    deduplicate,
    trim_whitespace,
    normalize_column,
)
from env.graders.task1_grader import grade_task1
from env.graders.task2_grader import grade_task2
from env.graders.task3_grader import grade_task3
from env.graders.final_evaluator import compute_final_score

TASK_SPECS = {
    1: {
        "id": "task1",
        "name": "Basic Data Cleaning",
        "difficulty": "easy",
        "objective": "Handle missing values, remove duplicates, fix simple type issues, and preserve dataset structure.",
        "grader": "env.graders.task1_grader.grade_task1",
    },
    2: {
        "id": "task2",
        "name": "Intermediate Data Cleaning",
        "difficulty": "medium",
        "objective": "Normalize numeric columns, reduce correlation issues, and preserve useful records.",
        "grader": "env.graders.task2_grader.grade_task2",
    },
    3: {
        "id": "task3",
        "name": "Full Data Cleaning Pipeline",
        "difficulty": "hard",
        "objective": "Perform end-to-end cleaning with null handling, duplicate removal, structure preservation, and minimal data loss.",
        "grader": "env.graders.task3_grader.grade_task3",
    },
}


class DataCleaningEnv:

    def __init__(self, task=1):
        self.task = int(task) if int(task) in TASK_SPECS else 1
        self.clean_df = None
        self.dirty_df = None
        self.manifest = None
        self.steps = 0
        self.done = False
        self.inspected_cols = set()

    def safe_df(self, df):
        return df.replace({float("nan"): None})

    def get_task_info(self):
        return TASK_SPECS[self.task]

    def reset(self):
        self.clean_df = generate_clean_dataset()
        self.dirty_df, self.manifest = inject_issues(self.clean_df)

        self.steps = 0
        self.done = False
        self.inspected_cols = set()

        return {
            "task": self.get_task_info(),
            "dataset": self.safe_df(self.dirty_df).to_dict(),
            "shape": list(self.dirty_df.shape),
            "steps": self.steps
        }

    def step(self, action):
        self.steps += 1
        reward = 0
        action_type = action.get("type")

        if action_type == "inspect_column":
            col = action["column"]
            if col not in self.inspected_cols:
                self.inspected_cols.add(col)
                reward += 0.01
            else:
                reward -= 0.02

        if action_type == "remove_nulls":
            col = action["column"]
            null_ratio = self.dirty_df[col].isnull().mean()
            if null_ratio > 0.3:
                self.dirty_df = remove_nulls(self.dirty_df, col)
                reward += 0.1
            else:
                reward -= 0.08

        elif action_type == "convert_types":
            self.dirty_df = convert_types(self.dirty_df, action["column"])
            reward += 0.1

        elif action_type == "deduplicate":
            self.dirty_df = deduplicate(self.dirty_df)
            reward += 0.1

        elif action_type == "trim_whitespace":
            self.dirty_df = trim_whitespace(self.dirty_df, action["column"])
            reward += 0.1

        elif action_type == "fill_nulls":
            col = action["column"]
            null_ratio = self.dirty_df[col].isnull().mean()
            if null_ratio < 0.3:
                self.dirty_df = fill_nulls(self.dirty_df, col)
                reward += 0.12
            else:
                reward -= 0.05

        elif action_type == "normalize":
            col = action["column"]
            if self.dirty_df[col].dtype != "object":
                self.dirty_df = normalize_column(self.dirty_df, col)
                reward += 0.1
            else:
                reward -= 0.08

        elif action_type == "drop_correlated":
            self.dirty_df = drop_correlated_feature(self.dirty_df, action["column"])
            reward += 0.1

        else:
            reward -= 0.05

        return {
            "task": self.get_task_info(),
            "dataset": self.safe_df(self.dirty_df).to_dict(),
            "shape": list(self.dirty_df.shape),
            "steps": self.steps
        }, reward, self.done, {}

    def state(self):
        return {
            "task": self.get_task_info(),
            "steps": self.steps,
            "dataset_shape": self.dirty_df.shape if self.dirty_df is not None else None,
            "inspected_columns": list(self.inspected_cols),
            "available_tasks": list(TASK_SPECS.values())
        }

    def submit_cleaned_data(self, agent_df):
        self.done = True

        if self.task == 1:
            quality = grade_task1(agent_df, self.clean_df, self.manifest)

        elif self.task == 2:
            quality = grade_task2(agent_df, self.clean_df)

        elif self.task == 3:
            quality = grade_task3(agent_df, self.clean_df)

        else:
            quality = 0.0

        final = compute_final_score(quality, self.steps)

        return {
            "task": self.get_task_info(),
            "quality_score": quality,
            "steps": self.steps,
            "final_score": final,
            "grader": self.get_task_info()["grader"]
        }

    def close(self):
        self.done = True


class StateManager:
    def __init__(self):
        self.steps = 0
        self.inspected_cols = set()