wanderlust-chatbot / app /training /split_dataset.py
Kiriten892's picture
feat(ML): add show_map intent + context-aware map routing
fda8066
Raw
History Blame Contribute Delete
5.42 kB
"""
Deterministic train/validation/test split for intent dataset.
Splits intent_train_enhanced.json into:
- intent_train.json (70%)
- intent_val.json (15%)
- intent_test.json (15%)
Stratified by intent to preserve class distribution across splits.
Fixed random seed (42) ensures reproducibility.
Usage:
python -m app.training.split_dataset [--input PATH] [--output-dir PATH]
"""
import json
import logging
import os
import sys
import argparse
import shutil
from collections import Counter
import numpy as np
from sklearn.model_selection import train_test_split
logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
logger = logging.getLogger(__name__)
SEED = 42
def load_dataset(path: str) -> list[dict]:
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
def save_dataset(data: list[dict], path: str) -> None:
os.makedirs(os.path.dirname(path), exist_ok=True)
with open(path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
logger.info(f"Saved {len(data)} samples -> {path}")
def split_dataset(
data: list[dict],
val_ratio: float = 0.15,
test_ratio: float = 0.15,
seed: int = SEED,
) -> tuple[list[dict], list[dict], list[dict]]:
"""Stratified split into train/val/test.
First carves out test set, then splits remaining into train/val.
Stratification ensures each intent is proportionally represented in all splits.
"""
texts = [item["text"] for item in data]
labels = [item["intent"] for item in data]
# First split: separate test set
train_val_data, test_data, train_val_labels, _ = train_test_split(
data,
labels,
test_size=test_ratio,
stratify=labels,
random_state=seed,
)
# Second split: train vs val from remaining
adjusted_val_ratio = val_ratio / (1.0 - test_ratio)
train_data, val_data, _, _ = train_test_split(
train_val_data,
train_val_labels,
test_size=adjusted_val_ratio,
stratify=train_val_labels,
random_state=seed,
)
return train_data, val_data, test_data
def print_split_summary(
train: list[dict],
val: list[dict],
test: list[dict],
) -> None:
total = len(train) + len(val) + len(test)
logger.info("=== Split Summary ===")
logger.info(
f"Train: {len(train)} ({len(train)/total*100:.1f}%) "
f"Val: {len(val)} ({len(val)/total*100:.1f}%) "
f"Test: {len(test)} ({len(test)/total*100:.1f}%)"
)
train_counts = Counter(item["intent"] for item in train)
val_counts = Counter(item["intent"] for item in val)
test_counts = Counter(item["intent"] for item in test)
all_intents = sorted(train_counts.keys())
print(f"\n{'Intent':<35} {'Train':>7} {'Val':>7} {'Test':>7} {'Total':>7}")
print("-" * 65)
for intent in all_intents:
t = train_counts.get(intent, 0)
v = val_counts.get(intent, 0)
s = test_counts.get(intent, 0)
print(f" {intent:<33} {t:>7} {v:>7} {s:>7} {t+v+s:>7}")
print("-" * 65)
print(f" {'TOTAL':<33} {len(train):>7} {len(val):>7} {len(test):>7} {total:>7}")
def main():
parser = argparse.ArgumentParser(description="Stratified train/val/test split")
parser.add_argument(
"--input",
default=None,
help="Input dataset JSON (default: intent_train_enhanced.json)",
)
parser.add_argument(
"--output-dir",
default=None,
help="Output directory (default: app/data/datasets/)",
)
parser.add_argument(
"--val-ratio", type=float, default=0.15, help="Validation split ratio (default: 0.15)"
)
parser.add_argument(
"--test-ratio", type=float, default=0.15, help="Test split ratio (default: 0.15)"
)
args = parser.parse_args()
base_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
datasets_dir = os.path.join(base_dir, "app", "data", "datasets")
input_path = args.input or os.path.join(datasets_dir, "intent_train_enhanced.json")
output_dir = args.output_dir or datasets_dir
train_out = os.path.join(output_dir, "intent_train.json")
val_out = os.path.join(output_dir, "intent_val.json")
test_out = os.path.join(output_dir, "intent_test.json")
backup_path = os.path.join(output_dir, "intent_train_v1_backup.json")
logger.info(f"Loading dataset: {input_path}")
data = load_dataset(input_path)
logger.info(f"Loaded {len(data)} samples across {len(set(d['intent'] for d in data))} intents")
# Back up existing intent_train.json before overwriting
if os.path.exists(train_out) and not os.path.exists(backup_path):
shutil.copy2(train_out, backup_path)
logger.info(f"Backed up existing intent_train.json -> {backup_path}")
elif os.path.exists(backup_path):
logger.info(f"Backup already exists at {backup_path}, skipping backup step")
logger.info(f"Splitting with seed={SEED}, val={args.val_ratio}, test={args.test_ratio}")
train_data, val_data, test_data = split_dataset(
data, val_ratio=args.val_ratio, test_ratio=args.test_ratio, seed=SEED
)
save_dataset(train_data, train_out)
save_dataset(val_data, val_out)
save_dataset(test_data, test_out)
print_split_summary(train_data, val_data, test_data)
if __name__ == "__main__":
main()