"""Synthetic grouped split demo; no biological measurements or model scores."""
import csv
import sys
from itertools import combinations

import numpy as np
from sklearn.model_selection import GroupShuffleSplit


def grouped_split(groups):
    groups = np.asarray(groups)
    if groups.ndim != 1 or len(set(groups.tolist())) < 5:
        raise ValueError("Supply a one-dimensional array with at least five groups")
    rows = np.arange(len(groups))
    first = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=17)
    dev, test = next(first.split(rows, groups=groups))
    second = GroupShuffleSplit(n_splits=1, test_size=0.25, random_state=29)
    train_local, valid_local = next(second.split(dev, groups=groups[dev]))
    # The second split returns positions within dev, not original row numbers.
    return {"train": dev[train_local], "valid": dev[valid_local], "test": test}


def verify(parts, groups, sequence_keys):
    rows = [int(i) for indices in parts.values() for i in indices]
    if sorted(rows) != list(range(len(groups))):
        raise ValueError("Every row must occur exactly once")
    if len(sequence_keys) != len(groups):
        raise ValueError("Sequence keys and groups must have equal length")
    for left, right in combinations(parts, 2):
        for values, label in [(groups, "group"), (sequence_keys, "exact sequence")]:
            overlap = {values[i] for i in parts[left]} & {values[i] for i in parts[right]}
            if overlap:
                raise ValueError(f"{left}/{right}: overlapping {label}")


def main():
    # Invented groups of unequal sizes; keys stand in for canonical sequences.
    groups = np.repeat(np.arange(10), np.arange(1, 11))
    keys = [f"synthetic-sequence-{i}" for i in range(len(groups))]
    parts = grouped_split(groups)
    verify(parts, groups, keys)
    writer = csv.writer(sys.stdout, lineterminator="\n")
    writer.writerow(["sample_id", "group_id", "split"])
    for name, indices in parts.items():
        group_count = len(set(groups[indices]))
        print(f"{name}: rows={len(indices)} groups={group_count}", file=sys.stderr)
        for i in indices:
            writer.writerow([f"sample_{i:03d}", int(groups[i]), name])
    print("PASS: rows, groups, and exact-sequence keys checked", file=sys.stderr)


if __name__ == "__main__":
    main()
