LeonardoMdSA/Context-aware-NLP-classification-platform-with-MCP
2
1"""2Seed and split dataset for training and evaluation.3 4- Reads: data/samples/training_data.json5- Writes:6 - data/samples/train.json7 - data/samples/eval.json8 9This script enforces:10- Stratified split by label11- Deterministic output (fixed random seed)12- Basic data validation13"""14 15import json16import random17from pathlib import Path18from collections import defaultdict19 20# -------------------------21# Configuration22# -------------------------23RANDOM_SEED = 4224TRAIN_RATIO = 0.725 26BASE_DIR = Path(__file__).resolve().parent.parent27SAMPLES_DIR = BASE_DIR / "data" / "samples"28 29SOURCE_FILE = SAMPLES_DIR / "training_data.json"30TRAIN_FILE = SAMPLES_DIR / "train.json"31EVAL_FILE = SAMPLES_DIR / "eval.json"32 33 34def main():35 if not SOURCE_FILE.exists():36 raise FileNotFoundError(f"Source dataset not found: {SOURCE_FILE}")37 38 with open(SOURCE_FILE, "r", encoding="utf-8") as f:39 data = json.load(f)40 41 if not isinstance(data, list) or len(data) == 0:42 raise ValueError("Dataset must be a non-empty list")43 44 # -------------------------45 # Basic validation46 # -------------------------47 for i, item in enumerate(data):48 if "text" not in item or "label" not in item:49 raise ValueError(f"Invalid sample at index {i}: {item}")50 51 # -------------------------52 # Stratified split53 # -------------------------54 random.seed(RANDOM_SEED)55 56 by_label = defaultdict(list)57 for item in data:58 by_label[item["label"]].append(item)59 60 train_data = []61 eval_data = []62 63 for label, items in by_label.items():64 random.shuffle(items)65 66 split_idx = max(1, int(len(items) * TRAIN_RATIO))67 68 train_data.extend(items[:split_idx])69 eval_data.extend(items[split_idx:])70 71 # Final shuffle (important)72 random.shuffle(train_data)73 random.shuffle(eval_data)74 75 # -------------------------76 # Write outputs77 # -------------------------78 SAMPLES_DIR.mkdir(parents=True, exist_ok=True)79 80 with open(TRAIN_FILE, "w", encoding="utf-8") as f:81 json.dump(train_data, f, indent=2, ensure_ascii=False)82 83 with open(EVAL_FILE, "w", encoding="utf-8") as f:84 json.dump(eval_data, f, indent=2, ensure_ascii=False)85 86 # -------------------------87 # Summary88 # -------------------------89 print("====================================")90 print("Dataset seeding completed")91 print("====================================")92 print(f"Total samples : {len(data)}")93 print(f"Train samples : {len(train_data)}")94 print(f"Eval samples : {len(eval_data)}")95 print()96 97 print("Label distribution (train):")98 _print_distribution(train_data)99 100 print("\nLabel distribution (eval):")101 _print_distribution(eval_data)102 103 104def _print_distribution(dataset):105 dist = defaultdict(int)106 for item in dataset:107 dist[item["label"]] += 1108 109 for label, count in sorted(dist.items()):110 print(f" {label:<20} {count}")111 112 113if __name__ == "__main__":114 main()115 