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.
- dataset_splitter/__init__.py +16 -0
- dataset_splitter/_io.py +32 -0
- dataset_splitter/_util.py +141 -0
- dataset_splitter/cli.py +104 -0
- dataset_splitter/groups.py +169 -0
- dataset_splitter/report.py +397 -0
- dataset_splitter/splitter.py +588 -0
- dataset_splitter-0.1.0.dist-info/METADATA +121 -0
- dataset_splitter-0.1.0.dist-info/RECORD +12 -0
- dataset_splitter-0.1.0.dist-info/WHEEL +4 -0
- dataset_splitter-0.1.0.dist-info/entry_points.txt +2 -0
- dataset_splitter-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -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
|
+
]
|
dataset_splitter/_io.py
ADDED
|
@@ -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
|
dataset_splitter/cli.py
ADDED
|
@@ -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
|