"""
Split dataset into train/validation sets
Stratified by section to ensure all sections represented
"""

import json
import random
from collections import defaultdict

def load_jsonl(filepath):
    """Load pretty-printed JSONL"""
    with open(filepath, 'r', encoding='utf-8') as f:
        content = f.read()

    conversations = []
    objects = content.split('\n}\n{')
    for i, obj in enumerate(objects):
        if i == 0:
            obj = obj + '\n}'
        else:
            obj = '{\n' + obj + '\n}'
        try:
            conversations.append(json.loads(obj))
        except:
            pass
    return conversations

def save_jsonl(conversations, filepath):
    """Save as JSONL"""
    with open(filepath, 'w', encoding='utf-8') as f:
        for conv in conversations:
            f.write(json.dumps(conv, ensure_ascii=False, indent=2) + '\n')

def stratified_split(conversations, val_ratio=0.1):
    """Split maintaining section representation"""
    # Group by section
    by_section = defaultdict(list)
    for conv in conversations:
        sec = conv['metadata']['section_code']
        by_section[sec].append(conv)

    train = []
    val = []

    for sec, convs in by_section.items():
        random.shuffle(convs)
        val_size = max(1, int(len(convs) * val_ratio))
        val.extend(convs[:val_size])
        train.extend(convs[val_size:])

    random.shuffle(train)
    random.shuffle(val)
    return train, val

def main():
    print("Loading dataset...")
    conversations = load_jsonl("gemma4_ayurveda_complete.jsonl")

    print(f"Total: {len(conversations)} conversations")

    # Stratified split
    train, val = stratified_split(conversations, val_ratio=0.1)

    print(f"Train: {len(train)}")
    print(f"Validation: {len(val)}")

    # Save splits
    save_jsonl(train, "gemma4_ayurveda_train.jsonl")
    save_jsonl(val, "gemma4_ayurveda_val.jsonl")

    # Save statistics
    stats = {
        "total": len(conversations),
        "train": len(train),
        "validation": len(val),
        "train_pct": len(train) / len(conversations) * 100,
        "val_pct": len(val) / len(conversations) * 100
    }

    with open("dataset_split_stats.json", 'w', encoding='utf-8') as f:
        json.dump(stats, f, indent=2)

    print("\nSaved:")
    print("  - gemma4_ayurveda_train.jsonl")
    print("  - gemma4_ayurveda_val.jsonl")
    print("  - dataset_split_stats.json")

if __name__ == "__main__":
    random.seed(42)
    main()
