EC256/openenv-data-engineering
0
1import os2import sys3 4# Ensure current directory is in path for imports5sys.path.insert(0, os.getcwd())6 7import json8from openai import OpenAI9 10from environment import DataEnv11from tasks import TASKS12 13 14def run_inference():15 print("[START] Initialization")16 17 # Required environment variables (strict format)18 API_BASE_URL = os.getenv("API_BASE_URL", "https://api.groq.com/openai/v1")19 MODEL_NAME = os.getenv("MODEL_NAME", "llama-3.3-70b-versatile")20 HF_TOKEN = os.getenv("HF_TOKEN") # No default - must be provided21 22 # Get API key from environment23 OPENAI_API_KEY = os.getenv("OPENAI_API_KEY", "")24 25 # Initialize LLM client using standard OpenAI client26 client = None27 if OPENAI_API_KEY:28 try:29 client = OpenAI(api_key=OPENAI_API_KEY, base_url=API_BASE_URL)30 print("[INFO] LLM client initialized")31 except Exception as e:32 print(f"[INFO] LLM client init failed: {e}")33 client = None34 35 env = DataEnv()36 results = {}37 38 for task_id in TASKS.keys():39 print(f"[STEP] Evaluating {task_id}")40 obs = env.reset(task_id=task_id)41 42 done = False43 step_count = 044 final_score = 0.045 46 while not done and step_count < 20:47 print(f"[STEP] Iterating (step {step_count + 1}) on {task_id}")48 49 prompt = f"""50You are a data engineering AI.51Dataset shape: {obs.total_rows} rows. Columns: {obs.columns}52Missing Values: {obs.missing_values}53Task: {obs.task_description}54 55Feedback from last action: {obs.feedback}56Debug Hints: {obs.debug_hints}57Sample Data: {json.dumps(obs.dataset_sample, default=str)}58 59Provide JSON representing your action. E.g., {{"action_type": "submit"}} or {{"action_type": "execute_pandas", "code": "df['column'] = ... "}}. Output raw JSON only.60"""61 62 # Use LLM if client available63 action_dict = {"action_type": "submit"}64 65 if OPENAI_API_KEY and client:66 try:67 response = client.chat.completions.create(68 model=MODEL_NAME,69 messages=[{"role": "user", "content": prompt}],70 response_format={"type": "json_object"},71 temperature=0.1,72 )73 action_dict = json.loads(response.choices[0].message.content)74 print(75 f"[STEP] LLM chose action: {action_dict.get('action_type', 'unknown')}"76 )77 except Exception as e:78 print(f"[STEP] LLM error: {e}")79 else:80 # Fallback heuristics when no LLM available81 if task_id == "easy_data_cleaning":82 if step_count == 0:83 action_dict = {84 "action_type": "execute_pandas",85 "code": "df['user_id'] = pd.to_numeric(df['user_id'].astype(str).str.replace('USR_', ''), errors='coerce').astype('Int64')",86 }87 elif step_count == 1:88 action_dict = {89 "action_type": "execute_pandas",90 "code": "df['signup_date'] = pd.to_datetime(df['signup_date'], errors='coerce', format='mixed').dt.strftime('%Y-%m-%d')",91 }92 elif step_count == 2:93 action_dict = {94 "action_type": "fill_nan",95 "column_name": "email",96 "fill_value": "unknown",97 }98 elif step_count >= 3:99 action_dict = {"action_type": "submit"}100 101 elif task_id == "medium_join_repair":102 if step_count == 0:103 action_dict = {104 "action_type": "execute_pandas",105 "code": "df['user_id'] = df['user_id'].astype(str).str.replace('USR_', '').astype('Int64')",106 }107 elif step_count == 1:108 action_dict = {109 "action_type": "merge_tables",110 "left_on": "user_id",111 "right_on": "user_id",112 "how": "left",113 }114 elif step_count == 2:115 action_dict = {116 "action_type": "parse_json",117 "column_name": "metadata",118 }119 elif step_count == 3:120 action_dict = {121 "action_type": "drop_column",122 "column_name": "metadata",123 }124 elif step_count >= 4:125 action_dict = {"action_type": "submit"}126 127 elif task_id == "hard_root_cause_analysis":128 if step_count == 0:129 action_dict = {130 "action_type": "execute_pandas",131 "code": "df['user_id'] = df['user_id'].astype(str).str.replace('USR_', '').astype('Int64')",132 }133 elif step_count == 1:134 action_dict = {135 "action_type": "merge_tables",136 "left_on": "user_id",137 "right_on": "user_id",138 "how": "left",139 }140 elif step_count == 2:141 action_dict = {142 "action_type": "parse_json",143 "column_name": "metadata",144 }145 elif step_count == 3:146 action_dict = {147 "action_type": "drop_column",148 "column_name": "metadata",149 }150 elif step_count == 4:151 action_dict = {152 "action_type": "execute_pandas",153 "code": "df['signup_date'] = pd.to_datetime(df['signup_date'], errors='coerce', format='mixed').dt.strftime('%Y-%m-%d')",154 }155 elif step_count == 5:156 action_dict = {157 "action_type": "fill_nan",158 "column_name": "email",159 "fill_value": "unknown",160 }161 elif step_count == 6:162 action_dict = {163 "action_type": "execute_pandas",164 "code": "df['amount'] = df['amount'].fillna(df['amount'].median())",165 }166 elif step_count >= 7:167 action_dict = {"action_type": "submit"}168 169 class DummyAction:170 def __init__(self, d):171 for k, v in d.items():172 setattr(self, k, v)173 if not hasattr(self, "action_type"):174 self.action_type = "submit"175 176 action = DummyAction(action_dict)177 obs, reward, done, info = env.step(action)178 final_score = reward.value179 180 step_count += 1181 182 results[task_id] = final_score183 print(f"[END] Task: {task_id} | Final Score: {final_score:.2f}")184 185 print("\n[SUMMARY] Results:")186 for tid, score in results.items():187 print(f" {tid}: {score:.2f}")188 avg = sum(results.values()) / len(results)189 print(f" Average: {avg:.2f}")190 191 192if __name__ == "__main__":193 run_inference()194 