dataset-splitter 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,16 @@
1
+ """dataset-splitter: leakage-safe train/validation/test splits in one call."""
2
+ from ._io import load_table
3
+ from .groups import detect_id_columns
4
+ from .report import SplitReport
5
+ from .splitter import Split, Splitter, split
6
+
7
+ __version__ = "0.1.0"
8
+ __all__ = [
9
+ "split",
10
+ "Splitter",
11
+ "Split",
12
+ "SplitReport",
13
+ "detect_id_columns",
14
+ "load_table",
15
+ "__version__",
16
+ ]
@@ -0,0 +1,32 @@
1
+ """Loading tabular input: a DataFrame is passed through, a path is read."""
2
+ from __future__ import annotations
3
+
4
+ import os
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ import pandas as pd
9
+
10
+
11
+ def load_table(data: Any) -> pd.DataFrame:
12
+ """Return `data` as a DataFrame; accepts a DataFrame or a path to .csv/.tsv/.parquet."""
13
+ if isinstance(data, pd.DataFrame):
14
+ return data
15
+ if isinstance(data, (str, os.PathLike)):
16
+ path = Path(data)
17
+ suffix = path.suffix.lower()
18
+ if suffix == ".csv":
19
+ return pd.read_csv(path)
20
+ if suffix == ".tsv":
21
+ return pd.read_csv(path, sep="\t")
22
+ if suffix in (".parquet", ".pq"):
23
+ try:
24
+ return pd.read_parquet(path)
25
+ except ImportError as exc:
26
+ raise ImportError(
27
+ "reading parquet needs pyarrow: pip install 'dataset-splitter[parquet]'"
28
+ ) from exc
29
+ raise ValueError(f"unsupported file type {suffix!r}; expected .csv, .tsv or .parquet")
30
+ raise TypeError(
31
+ f"expected a pandas DataFrame or a path to a .csv/.parquet file, got {type(data).__name__}"
32
+ )
@@ -0,0 +1,141 @@
1
+ """Small shared helpers: JSON conversion, time parsing, target kind and strata."""
2
+ from __future__ import annotations
3
+
4
+ import datetime as _dt
5
+ import math
6
+ import warnings
7
+ from typing import Any, List, Tuple
8
+
9
+ import numpy as np
10
+ import pandas as pd
11
+ from pandas.api.types import (
12
+ is_bool_dtype,
13
+ is_datetime64_any_dtype,
14
+ is_numeric_dtype,
15
+ is_object_dtype,
16
+ is_string_dtype,
17
+ )
18
+
19
+
20
+ def jsonable(obj: Any) -> Any:
21
+ """Recursively turn numpy/pandas scalars and containers into JSON-safe Python values."""
22
+ if obj is None or isinstance(obj, (bool, int, str)):
23
+ return obj
24
+ if isinstance(obj, float):
25
+ return obj if math.isfinite(obj) else None
26
+ if isinstance(obj, np.bool_):
27
+ return bool(obj)
28
+ if isinstance(obj, np.integer):
29
+ return int(obj)
30
+ if isinstance(obj, np.floating):
31
+ value = float(obj)
32
+ return value if math.isfinite(value) else None
33
+ if isinstance(obj, dict):
34
+ return {str(k): jsonable(v) for k, v in obj.items()}
35
+ if isinstance(obj, (list, tuple, set, frozenset, np.ndarray, pd.Index, pd.Series)):
36
+ return [jsonable(v) for v in list(obj)]
37
+ if obj is pd.NaT:
38
+ return None
39
+ if isinstance(obj, np.datetime64):
40
+ return None if np.isnat(obj) else pd.Timestamp(obj).isoformat()
41
+ if isinstance(obj, (pd.Timestamp, _dt.datetime, _dt.date)):
42
+ return obj.isoformat()
43
+ if isinstance(obj, (pd.Timedelta, _dt.timedelta, pd.Interval)):
44
+ return str(obj)
45
+ try:
46
+ if pd.isna(obj):
47
+ return None
48
+ except (TypeError, ValueError):
49
+ pass
50
+ return str(obj)
51
+
52
+
53
+ def parse_time(series: pd.Series, name: Any) -> pd.Series:
54
+ """Return a time column as datetime64 or float64 values.
55
+
56
+ Raises ValueError when values are missing or cannot be interpreted as dates or numbers.
57
+ """
58
+ if is_datetime64_any_dtype(series):
59
+ parsed = series
60
+ elif is_bool_dtype(series):
61
+ raise ValueError(f"time column {name!r} is boolean; expected dates or numbers")
62
+ elif is_numeric_dtype(series):
63
+ parsed = pd.to_numeric(series, errors="coerce").astype("float64")
64
+ else:
65
+ parsed = _to_datetime(series, name)
66
+ n_missing = int(parsed.isna().sum())
67
+ if n_missing:
68
+ raise ValueError(
69
+ f"time column {name!r} has {n_missing} missing value(s); fill or drop them before splitting"
70
+ )
71
+ if is_datetime64_any_dtype(parsed) and getattr(parsed.dt, "tz", None) is not None:
72
+ parsed = parsed.dt.tz_convert("UTC").dt.tz_localize(None)
73
+ return parsed
74
+
75
+
76
+ def _to_datetime(series: pd.Series, name: Any) -> pd.Series:
77
+ values = series.astype(object) if not is_object_dtype(series) else series
78
+ with warnings.catch_warnings():
79
+ warnings.simplefilter("ignore")
80
+ try:
81
+ return pd.to_datetime(values, errors="raise")
82
+ except (ValueError, TypeError, OverflowError):
83
+ pass
84
+ try: # pandas >= 2.0 can parse per-element formats
85
+ return pd.to_datetime(values, errors="raise", format="mixed")
86
+ except (ValueError, TypeError, OverflowError) as exc:
87
+ raise ValueError(
88
+ f"time column {name!r} could not be parsed as dates or numbers"
89
+ ) from exc
90
+
91
+
92
+ def time_key(parsed: pd.Series) -> np.ndarray:
93
+ """A sortable numeric array for a column returned by parse_time()."""
94
+ if is_datetime64_any_dtype(parsed):
95
+ return parsed.to_numpy(dtype="datetime64[ns]").astype("int64")
96
+ return parsed.to_numpy(dtype="float64")
97
+
98
+
99
+ def target_kind(y: pd.Series, max_categories: int = 20) -> str:
100
+ """'categorical' for non-numeric or low-cardinality targets, otherwise 'numeric'."""
101
+ if (
102
+ isinstance(y.dtype, pd.CategoricalDtype)
103
+ or is_bool_dtype(y)
104
+ or is_object_dtype(y)
105
+ or is_string_dtype(y)
106
+ ):
107
+ return "categorical"
108
+ if is_datetime64_any_dtype(y):
109
+ return "numeric"
110
+ if is_numeric_dtype(y):
111
+ return "categorical" if y.nunique(dropna=True) <= max_categories else "numeric"
112
+ return "categorical"
113
+
114
+
115
+ def strata(y: pd.Series, kind: str, n_bins: int) -> Tuple[np.ndarray, List[str]]:
116
+ """Integer stratum code per row plus the label of each code.
117
+
118
+ Categorical targets use their classes; numeric targets use quantile bins.
119
+ Missing target values become a stratum of their own, labelled '<missing>'.
120
+ """
121
+ if kind == "categorical":
122
+ codes, uniques = pd.factorize(y)
123
+ labels = [str(u) for u in uniques]
124
+ else:
125
+ categories = None
126
+ try:
127
+ binned = pd.qcut(y, q=max(1, int(n_bins)), duplicates="drop")
128
+ categories = list(binned.cat.categories)
129
+ except (ValueError, TypeError):
130
+ binned = None
131
+ if binned is None or not categories:
132
+ codes = np.zeros(len(y), dtype=np.int64)
133
+ labels = ["all"]
134
+ else:
135
+ codes = binned.cat.codes.to_numpy()
136
+ labels = [str(iv) for iv in categories]
137
+ codes = np.asarray(codes, dtype=np.int64)
138
+ if (codes < 0).any():
139
+ codes = np.where(codes < 0, len(labels), codes)
140
+ labels = labels + ["<missing>"]
141
+ return codes, labels
@@ -0,0 +1,104 @@
1
+ """Command line entry point: ``dataset-splitter data.csv --target y --group customer_id``."""
2
+ from __future__ import annotations
3
+
4
+ import argparse
5
+ import json
6
+ import sys
7
+ from typing import List, Optional, Union
8
+
9
+ from . import __version__
10
+ from .splitter import Splitter
11
+
12
+
13
+ def _size(text: str) -> Union[int, float]:
14
+ """'0.2' -> fraction, '100' -> absolute row count."""
15
+ if "." in text or "e" in text.lower():
16
+ return float(text)
17
+ return int(text)
18
+
19
+
20
+ def build_parser() -> argparse.ArgumentParser:
21
+ parser = argparse.ArgumentParser(
22
+ prog="dataset-splitter",
23
+ description=(
24
+ "Leakage-safe train/validation/test splits: stratified, grouped, time-aware, "
25
+ "and checked. Prints the split report; exit status 1 when the report is not ok."
26
+ ),
27
+ )
28
+ parser.add_argument("path", help="input table (.csv, .tsv or .parquet)")
29
+ parser.add_argument(
30
+ "--target", metavar="COL", help="column to stratify on and report class balance for"
31
+ )
32
+ parser.add_argument(
33
+ "--group",
34
+ metavar="COL",
35
+ nargs="+",
36
+ help="column(s) whose rows must stay on one side, or 'auto' to detect id-like columns",
37
+ )
38
+ parser.add_argument(
39
+ "--time", metavar="COL", help="chronological split on this column (oldest train, newest test)"
40
+ )
41
+ parser.add_argument(
42
+ "--test-size", type=_size, default=0.2, help="fraction or row count (default 0.2)"
43
+ )
44
+ parser.add_argument(
45
+ "--val-size", type=_size, default=0.1, help="fraction or row count (default 0.1)"
46
+ )
47
+ parser.add_argument("--random-state", type=int, default=0, help="seed for shuffling (default 0)")
48
+ parser.add_argument(
49
+ "--no-dedupe", action="store_true", help="do not keep exact duplicate rows on the same side"
50
+ )
51
+ parser.add_argument("--json", action="store_true", help="print the report as JSON")
52
+ parser.add_argument(
53
+ "--output", metavar="DIR", help="write train/val/test files and report.json here"
54
+ )
55
+ parser.add_argument(
56
+ "--format", choices=("csv", "parquet"), default="csv", help="output file format"
57
+ )
58
+ parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
59
+ return parser
60
+
61
+
62
+ def main(argv: Optional[List[str]] = None) -> int:
63
+ """Parse arguments, split the table, print the report. Returns the process exit status."""
64
+ for stream in (sys.stdout, sys.stderr):
65
+ if hasattr(stream, "reconfigure"):
66
+ try:
67
+ stream.reconfigure(encoding="utf-8", errors="replace")
68
+ except (ValueError, OSError): # a stream that cannot be reconfigured
69
+ pass
70
+ args = build_parser().parse_args(argv)
71
+ group = None
72
+ if args.group:
73
+ if args.group == ["auto"]:
74
+ group = "auto"
75
+ elif len(args.group) == 1:
76
+ group = args.group[0]
77
+ else:
78
+ group = list(args.group)
79
+ try:
80
+ splitter = Splitter(
81
+ target=args.target,
82
+ group=group,
83
+ time=args.time,
84
+ test_size=args.test_size,
85
+ val_size=args.val_size,
86
+ random_state=args.random_state,
87
+ dedupe=not args.no_dedupe,
88
+ )
89
+ result = splitter.split(args.path)
90
+ report = result.report()
91
+ if args.output:
92
+ result.save(args.output, format=args.format)
93
+ except (ValueError, TypeError, ImportError, FileNotFoundError) as exc:
94
+ print(f"error: {exc}", file=sys.stderr)
95
+ return 2
96
+ if args.json:
97
+ print(json.dumps(report.to_dict(), indent=2, ensure_ascii=False))
98
+ else:
99
+ print(report.summary())
100
+ return 0 if report.ok else 1
101
+
102
+
103
+ if __name__ == "__main__": # pragma: no cover
104
+ sys.exit(main())
@@ -0,0 +1,169 @@
1
+ """Row grouping: id-like column detection, composite keys, linked ids, exact duplicates.
2
+
3
+ Every function here returns a compact integer id per row (0..k-1). Rows that share an id
4
+ must land on the same side of a split.
5
+ """
6
+ from __future__ import annotations
7
+
8
+ import re
9
+ from typing import Any, Iterable, List, Optional, Sequence
10
+
11
+ import numpy as np
12
+ import pandas as pd
13
+ from pandas.api.types import is_float_dtype
14
+
15
+ _ID_SUFFIXES = {"id", "ids", "uuid", "guid"}
16
+ _ID_ANYWHERE = {"uuid", "guid"}
17
+ _ENTITY_WORDS = {
18
+ "customer", "user", "session", "subject", "patient", "account", "device", "household",
19
+ "client", "member", "visitor", "participant", "person", "student", "employee", "player",
20
+ "entity", "case", "tenant", "organisation", "organization", "org", "company", "merchant",
21
+ "driver", "vehicle", "site", "store", "clinic", "hospital", "school", "family",
22
+ }
23
+ _NUMBER_WORDS = {"no", "num", "number", "nbr", "code", "key", "ref"}
24
+ _COMPOUNDS = {w + "id" for w in _ENTITY_WORDS} | {
25
+ "orderid", "itemid", "productid", "docid", "fileid", "groupid", "recordid", "txid",
26
+ "transactionid", "eventid", "requestid", "traceid",
27
+ }
28
+ _SPLIT_RE = re.compile(r"[^0-9a-zA-Z]+")
29
+ _CAMEL_RE = re.compile(r"(?<=[a-z0-9])(?=[A-Z])")
30
+
31
+
32
+ def _tokens(name: Any) -> List[str]:
33
+ text = _CAMEL_RE.sub("_", str(name))
34
+ return [t for t in _SPLIT_RE.split(text.lower()) if t]
35
+
36
+
37
+ def is_id_like(name: Any) -> bool:
38
+ """True when a column name looks like an entity identifier (customer_id, userId, uuid...)."""
39
+ tokens = _tokens(name)
40
+ if not tokens:
41
+ return False
42
+ if tokens[-1] in _ID_SUFFIXES:
43
+ return True
44
+ if any(t in _ID_ANYWHERE for t in tokens):
45
+ return True
46
+ if "".join(tokens) in _COMPOUNDS:
47
+ return True
48
+ if len(tokens) == 1 and tokens[0] in _ENTITY_WORDS:
49
+ return True
50
+ if tokens[-1] in _NUMBER_WORDS and any(t in _ENTITY_WORDS for t in tokens[:-1]):
51
+ return True
52
+ return False
53
+
54
+
55
+ def detect_id_columns(df: pd.DataFrame, exclude: Optional[Iterable[Any]] = None) -> List[Any]:
56
+ """Columns that look like ids by name and whose values repeat (so grouping matters).
57
+
58
+ Constant columns and columns unique per row are skipped, as are `exclude`d columns.
59
+ """
60
+ skip = set(exclude or ())
61
+ n = len(df)
62
+ found: List[Any] = []
63
+ for col in df.columns:
64
+ if col in skip or not is_id_like(col):
65
+ continue
66
+ series = df[col]
67
+ if isinstance(series, pd.DataFrame): # duplicated column name
68
+ continue
69
+ if is_float_dtype(series):
70
+ values = series.dropna()
71
+ if len(values) and not np.all(np.mod(values.to_numpy(dtype="float64"), 1) == 0):
72
+ continue # fractional floats are measurements, not ids
73
+ try:
74
+ n_unique = int(series.nunique(dropna=True))
75
+ except TypeError:
76
+ continue
77
+ if n_unique < 2 or n_unique >= n:
78
+ continue
79
+ found.append(col)
80
+ return found
81
+
82
+
83
+ def _codes(df: pd.DataFrame, col: Any) -> np.ndarray:
84
+ series = df[col]
85
+ if isinstance(series, pd.DataFrame):
86
+ raise ValueError(f"column {col!r} appears more than once in the frame")
87
+ try:
88
+ codes, _ = pd.factorize(series)
89
+ except TypeError as exc:
90
+ raise ValueError(f"group column {col!r} contains unhashable values") from exc
91
+ return np.asarray(codes, dtype=np.int64)
92
+
93
+
94
+ def _isolate_missing(ids: np.ndarray, missing: np.ndarray) -> np.ndarray:
95
+ """Give rows with a missing key an id of their own, then renumber compactly."""
96
+ ids = ids.copy()
97
+ if missing.any():
98
+ start = int(ids[~missing].max()) + 1 if (~missing).any() else 0
99
+ ids[missing] = start + np.arange(int(missing.sum()))
100
+ return np.asarray(pd.factorize(ids)[0], dtype=np.int64)
101
+
102
+
103
+ def composite_ids(df: pd.DataFrame, cols: Sequence[Any]) -> np.ndarray:
104
+ """One id per distinct combination of `cols`; rows missing any key stand alone."""
105
+ n = len(df)
106
+ if n == 0:
107
+ return np.zeros(0, dtype=np.int64)
108
+ code_list = [_codes(df, c) for c in cols]
109
+ missing = np.zeros(n, dtype=bool)
110
+ for codes in code_list:
111
+ missing |= codes < 0
112
+ if len(code_list) == 1:
113
+ ids = code_list[0]
114
+ else:
115
+ ids = np.zeros(n, dtype=np.int64)
116
+ for codes in code_list:
117
+ width = int(codes.max()) + 2
118
+ ids = np.asarray(pd.factorize(ids * width + (codes + 1))[0], dtype=np.int64)
119
+ return _isolate_missing(ids, missing)
120
+
121
+
122
+ def link_codes(code_list: Sequence[np.ndarray]) -> np.ndarray:
123
+ """Connected components over rows: two rows are linked when they share a value in any array.
124
+
125
+ Negative codes mean "no value" and never link anything.
126
+ """
127
+ n = len(code_list[0])
128
+ labels = np.arange(n, dtype=np.int64)
129
+ while True:
130
+ changed = False
131
+ for codes in code_list:
132
+ valid = codes >= 0
133
+ if not valid.any():
134
+ continue
135
+ current = labels[valid]
136
+ lowest = (
137
+ pd.Series(current).groupby(codes[valid], sort=False).transform("min").to_numpy()
138
+ )
139
+ if np.any(lowest != current):
140
+ labels[valid] = lowest
141
+ changed = True
142
+ if not changed:
143
+ break
144
+ return np.asarray(pd.factorize(labels)[0], dtype=np.int64)
145
+
146
+
147
+ def linked_ids(df: pd.DataFrame, cols: Sequence[Any]) -> np.ndarray:
148
+ """One id per connected component: rows sharing a value in any of `cols` stay together."""
149
+ if len(df) == 0:
150
+ return np.zeros(0, dtype=np.int64)
151
+ return link_codes([_codes(df, c) for c in cols])
152
+
153
+
154
+ def duplicate_ids(df: pd.DataFrame) -> np.ndarray:
155
+ """One id per distinct row (all columns compared); exact duplicates share an id."""
156
+ n = len(df)
157
+ if n == 0 or df.shape[1] == 0:
158
+ return np.arange(n, dtype=np.int64)
159
+ ids = np.zeros(n, dtype=np.int64)
160
+ for i in range(df.shape[1]):
161
+ column = df.iloc[:, i]
162
+ try:
163
+ codes = pd.factorize(column)[0]
164
+ except TypeError: # unhashable cells (lists, dicts): compare their text form
165
+ codes = pd.factorize(column.astype(str))[0]
166
+ codes = np.asarray(codes, dtype=np.int64) + 1 # missing (-1) becomes a value of its own
167
+ width = int(codes.max()) + 1
168
+ ids = np.asarray(pd.factorize(ids * width + codes)[0], dtype=np.int64)
169
+ return ids