Naphula commited on
Commit
957e88b
ยท
verified ยท
1 Parent(s): bf21b97

Upload dataset_patcher.py

Browse files
Files changed (1) hide show
  1. dataset_patcher.py +346 -0
dataset_patcher.py ADDED
@@ -0,0 +1,346 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+ import sys
4
+ from collections import Counter
5
+ from typing import Any, Dict, List, Optional, Tuple
6
+
7
+ try:
8
+ import pandas as pd
9
+ import numpy as np
10
+ except ImportError:
11
+ pd = None
12
+ np = None
13
+
14
+ USER_ROLES = {"human", "user", "prompter"}
15
+ ASSISTANT_ROLES = {"gpt", "assistant"}
16
+ SYSTEM_ROLES = {"system"}
17
+
18
+
19
+ def to_python_list(val: Any) -> Any:
20
+ """Safely converts numpy arrays / pandas series / iterables to native python list without truth checks."""
21
+ if val is None:
22
+ return None
23
+ if isinstance(val, (list, tuple)):
24
+ return list(val)
25
+ if np is not None and isinstance(val, np.ndarray):
26
+ return val.tolist()
27
+ if hasattr(val, "tolist"):
28
+ return val.tolist()
29
+ return val
30
+
31
+
32
+ def inspect_entry(
33
+ entry: Dict[str, Any],
34
+ ) -> Tuple[bool, List[str], Optional[List[Dict[str, str]]]]:
35
+ errors = []
36
+
37
+ # Safely retrieve conversations without using `or` operator across arrays
38
+ raw_conv = entry.get("conversations")
39
+ if raw_conv is None or (isinstance(raw_conv, float) and pd.isna(raw_conv)):
40
+ raw_conv = entry.get("messages")
41
+
42
+ conv = to_python_list(raw_conv)
43
+
44
+ # Handle stringified JSON inside parquet columns
45
+ if isinstance(conv, str):
46
+ try:
47
+ conv = json.loads(conv)
48
+ conv = to_python_list(conv)
49
+ except Exception:
50
+ return (
51
+ False,
52
+ ["Field 'conversations'/'messages' contains invalid JSON string."],
53
+ None,
54
+ )
55
+
56
+ if conv is None:
57
+ return False, ["Missing 'conversations' or 'messages' key."], None
58
+ if not isinstance(conv, list):
59
+ return False, ["Conversation is not a list/array."], None
60
+ if len(conv) == 0:
61
+ return False, ["Conversation list is empty."], None
62
+
63
+ normalized_turns = []
64
+
65
+ for idx, turn in enumerate(conv):
66
+ # Handle stringified sub-elements if any
67
+ if isinstance(turn, str):
68
+ try:
69
+ turn = json.loads(turn)
70
+ except Exception:
71
+ errors.append(f"Turn #{idx} is not a valid dictionary or JSON.")
72
+ continue
73
+
74
+ if not isinstance(turn, dict):
75
+ errors.append(f"Turn #{idx} is not an object/dictionary.")
76
+ continue
77
+
78
+ raw_role = str(turn.get("from") or turn.get("role") or "").strip()
79
+ raw_content = str(turn.get("value") or turn.get("content") or "").strip()
80
+
81
+ if not raw_role:
82
+ errors.append(f"Turn #{idx} is missing role identifier.")
83
+ if not raw_content:
84
+ errors.append(f"Turn #{idx} has empty text content.")
85
+
86
+ role_lower = raw_role.lower()
87
+ canonical_role = None
88
+ if role_lower in USER_ROLES:
89
+ canonical_role = "human"
90
+ elif role_lower in ASSISTANT_ROLES:
91
+ canonical_role = "gpt"
92
+ elif role_lower in SYSTEM_ROLES:
93
+ canonical_role = "system"
94
+ else:
95
+ errors.append(f"Turn #{idx} has unrecognized role: '{raw_role}'.")
96
+
97
+ role_key = "from" if "from" in turn else "role"
98
+ val_key = "value" if "value" in turn else "content"
99
+
100
+ if raw_content and canonical_role:
101
+ normalized_turns.append(
102
+ {role_key: canonical_role, val_key: raw_content}
103
+ )
104
+
105
+ if not normalized_turns:
106
+ return False, errors or ["No valid text turns found."], None
107
+
108
+ last_turn_role = (
109
+ normalized_turns[-1].get("from") or normalized_turns[-1].get("role")
110
+ )
111
+ if last_turn_role != "gpt":
112
+ errors.append(f"Last turn is '{last_turn_role}' (must end with 'gpt').")
113
+
114
+ has_human = any(
115
+ (t.get("from") or t.get("role")) == "human" for t in normalized_turns
116
+ )
117
+ has_gpt = any(
118
+ (t.get("from") or t.get("role")) == "gpt" for t in normalized_turns
119
+ )
120
+
121
+ if not has_human:
122
+ errors.append("Missing at least one user/human turn.")
123
+ if not has_gpt:
124
+ errors.append("Missing at least one assistant/gpt turn.")
125
+
126
+ consecutive_dupes = False
127
+ for i in range(len(normalized_turns) - 1):
128
+ r1 = (
129
+ normalized_turns[i].get("from") or normalized_turns[i].get("role")
130
+ )
131
+ r2 = (
132
+ normalized_turns[i + 1].get("from")
133
+ or normalized_turns[i + 1].get("role")
134
+ )
135
+ if r1 == r2 and r1 != "system":
136
+ consecutive_dupes = True
137
+ break
138
+ if consecutive_dupes:
139
+ errors.append("Contains consecutive turns with the same role.")
140
+
141
+ repaired = None
142
+ if errors:
143
+ repaired = repair_conversation(normalized_turns)
144
+
145
+ is_valid = len(errors) == 0
146
+ return is_valid, errors, repaired
147
+
148
+
149
+ def repair_conversation(
150
+ turns: List[Dict[str, str]],
151
+ ) -> Optional[List[Dict[str, str]]]:
152
+ if not turns:
153
+ return None
154
+
155
+ merged: List[Dict[str, str]] = []
156
+ for turn in turns:
157
+ role_k = "from" if "from" in turn else "role"
158
+ val_k = "value" if "value" in turn else "content"
159
+ curr_role = turn[role_k]
160
+ curr_val = turn[val_k]
161
+
162
+ if (
163
+ merged
164
+ and (merged[-1].get("from") or merged[-1].get("role")) == curr_role
165
+ ):
166
+ prev_val_k = "value" if "value" in merged[-1] else "content"
167
+ merged[-1][prev_val_k] += f"\n\n{curr_val}"
168
+ else:
169
+ merged.append({role_k: curr_role, val_k: curr_val})
170
+
171
+ while (
172
+ merged and (merged[-1].get("from") or merged[-1].get("role")) != "gpt"
173
+ ):
174
+ merged.pop()
175
+
176
+ has_human = any(
177
+ (t.get("from") or t.get("role")) == "human" for t in merged
178
+ )
179
+ has_gpt = any((t.get("from") or t.get("role")) == "gpt" for t in merged)
180
+
181
+ if merged and has_human and has_gpt:
182
+ return merged
183
+ return None
184
+
185
+
186
+ def load_file(path: str) -> Tuple[List[Dict[str, Any]], str]:
187
+ ext = os.path.splitext(path)[1].lower()
188
+
189
+ if ext == ".parquet":
190
+ if pd is None:
191
+ raise ImportError(
192
+ "Reading .parquet requires pandas and pyarrow. Run: pip install pandas pyarrow"
193
+ )
194
+ df = pd.read_parquet(path)
195
+ return df.to_dict(orient="records"), "parquet"
196
+
197
+ with open(path, "r", encoding="utf-8") as f:
198
+ first_char = f.read(1)
199
+ f.seek(0)
200
+ if first_char == "[":
201
+ return json.load(f), "json"
202
+ else:
203
+ records = []
204
+ for line in f:
205
+ line = line.strip()
206
+ if line:
207
+ records.append(json.loads(line))
208
+ return records, "jsonl"
209
+
210
+
211
+ def main():
212
+ if len(sys.argv) < 2:
213
+ input_path = input(
214
+ "Enter path to JSON, JSONL, or Parquet dataset: "
215
+ ).strip()
216
+ else:
217
+ input_path = sys.argv[1]
218
+
219
+ input_path = input_path.strip("\"'")
220
+
221
+ if not os.path.isfile(input_path):
222
+ print(f"File not found: {input_path}")
223
+ sys.exit(1)
224
+
225
+ print(f"\n๐Ÿ“‚ Loading {input_path} into memory...")
226
+ try:
227
+ records, _ = load_file(input_path)
228
+ except Exception as e:
229
+ print(f"Error reading file: {e}")
230
+ sys.exit(1)
231
+
232
+ total_count = len(records)
233
+ print(f"โœ… Loaded {total_count} records. Starting audit...\n")
234
+
235
+ valid_indices = []
236
+ invalid_data = []
237
+ error_counter = Counter()
238
+
239
+ for idx, record in enumerate(records):
240
+ is_valid, errors, repaired = inspect_entry(record)
241
+ if is_valid:
242
+ valid_indices.append(idx)
243
+ else:
244
+ invalid_data.append((idx, record, errors, repaired))
245
+ for err in errors:
246
+ error_counter[err] += 1
247
+
248
+ invalid_count = len(invalid_data)
249
+ valid_count = len(valid_indices)
250
+
251
+ print("=" * 60)
252
+ print("๐Ÿ“Š AUDIT RESULTS SUMMARY")
253
+ print("=" * 60)
254
+ print(f"Total records audited : {total_count}")
255
+ print(f"Clean records : {valid_count} ({(valid_count/total_count)*100:.2f}%)")
256
+ print(f"Malformed records : {invalid_count} ({(invalid_count/total_count)*100:.2f}%)")
257
+ print("\nError Breakdown:")
258
+ for err, cnt in error_counter.most_common():
259
+ print(f" โ€ข [{cnt} occurrences] {err}")
260
+ print("=" * 60)
261
+
262
+ if invalid_count == 0:
263
+ print("\nโœจ No issues detected. Your dataset is 100% compliant.")
264
+ return
265
+
266
+ preview_limit = min(3, invalid_count)
267
+ print(f"\n๐Ÿ” PREVIEWING FIRST {preview_limit} MALFORMED ENTRIES:")
268
+ for i in range(preview_limit):
269
+ orig_idx, rec, errs, rep = invalid_data[i]
270
+ print(f"\n--- [Record #{orig_idx}] ---")
271
+ print(f"Issues Detected: {errs}")
272
+ raw_conv = rec.get("conversations")
273
+ if raw_conv is None:
274
+ raw_conv = rec.get("messages")
275
+ conv_preview = str(to_python_list(raw_conv))[:250]
276
+ print(f"Content: {conv_preview}...")
277
+ print(f"Repairable: {'Yes' if rep is not None else 'No'}")
278
+
279
+ print("\n" + "=" * 60)
280
+ print("๐Ÿ› ๏ธ RESOLUTION OPTIONS:")
281
+ print(" [1] DELETE malformed records (keep only the 100% clean ones).")
282
+ print(" [2] REPAIR what is recoverable (drop only unfixable entries).")
283
+ print(" [3] CANCEL and make no changes.")
284
+ print("=" * 60)
285
+
286
+ choice = ""
287
+ while choice not in ["1", "2", "3"]:
288
+ choice = input("Select an option [1/2/3]: ").strip()
289
+
290
+ if choice == "3":
291
+ print("\nAborted. No changes written.")
292
+ sys.exit(0)
293
+
294
+ output_records = []
295
+ if choice == "1":
296
+ output_records = [records[i] for i in valid_indices]
297
+ elif choice == "2":
298
+ output_records = [records[i] for i in valid_indices]
299
+ repaired_success = 0
300
+ unrepairable_dropped = 0
301
+ for orig_idx, rec, errs, rep in invalid_data:
302
+ if rep is not None:
303
+ target_key = (
304
+ "conversations" if "conversations" in rec else "messages"
305
+ )
306
+ rec[target_key] = rep
307
+ output_records.append(rec)
308
+ repaired_success += 1
309
+ else:
310
+ unrepairable_dropped += 1
311
+ print(f"\n โ€ข Successfully repaired : {repaired_success}")
312
+ print(f" โ€ข Unrepairable & dropped: {unrepairable_dropped}")
313
+
314
+ # Standardize output: ensure all array elements inside each record are native Python lists
315
+ for rec in output_records:
316
+ for k, v in list(rec.items()):
317
+ rec[k] = to_python_list(v)
318
+
319
+ base = os.path.splitext(input_path)[0]
320
+ default_out = f"{base}_clean.jsonl"
321
+ out_path = input(
322
+ f"\nEnter output path [Default: {default_out}]: "
323
+ ).strip().strip("\"'")
324
+ if not out_path:
325
+ out_path = default_out
326
+
327
+ print(f"\n๐Ÿ’พ Writing {len(output_records)} records to {out_path}...")
328
+ out_ext = os.path.splitext(out_path)[1].lower()
329
+
330
+ if out_ext == ".parquet":
331
+ df_out = pd.DataFrame(output_records)
332
+ df_out.to_parquet(out_path, index=False)
333
+ elif out_ext == ".json":
334
+ with open(out_path, "w", encoding="utf-8") as f:
335
+ json.dump(output_records, f, ensure_ascii=False, indent=2)
336
+ else:
337
+ # Default to jsonl
338
+ with open(out_path, "w", encoding="utf-8") as f:
339
+ for row in output_records:
340
+ f.write(json.dumps(row, ensure_ascii=False) + "\n")
341
+
342
+ print(f"โœ… Finished! Saved to: {os.path.abspath(out_path)}\n")
343
+
344
+
345
+ if __name__ == "__main__":
346
+ main()