easyclassifier 0.8.1__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.
- easyclassifier/__init__.py +23 -0
- easyclassifier/__main__.py +61 -0
- easyclassifier/dataset.py +142 -0
- easyclassifier/demo_data.py +32 -0
- easyclassifier/diagnostics.py +95 -0
- easyclassifier/distances.py +152 -0
- easyclassifier/evaluation.py +340 -0
- easyclassifier/figures.py +510 -0
- easyclassifier/help_texts.py +113 -0
- easyclassifier/importance.py +79 -0
- easyclassifier/latex_report.py +555 -0
- easyclassifier/logbook.py +31 -0
- easyclassifier/models.py +94 -0
- easyclassifier/preprocessing.py +219 -0
- easyclassifier/recommend.py +128 -0
- easyclassifier/reporting.py +70 -0
- easyclassifier/target.py +202 -0
- easyclassifier/ui.py +281 -0
- easyclassifier/wizard.py +1266 -0
- easyclassifier-0.8.1.dist-info/METADATA +267 -0
- easyclassifier-0.8.1.dist-info/RECORD +25 -0
- easyclassifier-0.8.1.dist-info/WHEEL +5 -0
- easyclassifier-0.8.1.dist-info/entry_points.txt +2 -0
- easyclassifier-0.8.1.dist-info/licenses/LICENSE +21 -0
- easyclassifier-0.8.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""EasyClassifier - Machine Learning without Programming.
|
|
2
|
+
|
|
3
|
+
A guided, menu-driven wizard that lets non-programmers build and evaluate
|
|
4
|
+
classification models from a CSV file.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
__version__ = "0.8.1"
|
|
8
|
+
__author__ = "Ahmad Hassanat"
|
|
9
|
+
|
|
10
|
+
CITATION = (
|
|
11
|
+
"Hassanat, A. (2026). EasyClassifier: Machine Learning without "
|
|
12
|
+
"Programming (Version {version}) [Software].".format(version=__version__)
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
from .wizard import Wizard # noqa: E402
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def run():
|
|
19
|
+
"""Launch the interactive wizard."""
|
|
20
|
+
Wizard().run()
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
__all__ = ["Wizard", "run", "__version__", "CITATION"]
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
"""Console entry point for EasyClassifier.
|
|
2
|
+
|
|
3
|
+
Start the wizard with either of:
|
|
4
|
+
|
|
5
|
+
easyclassifier
|
|
6
|
+
python -m easyclassifier
|
|
7
|
+
|
|
8
|
+
Options:
|
|
9
|
+
|
|
10
|
+
--version show the installed version and exit
|
|
11
|
+
--help show this help and exit
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
import sys
|
|
15
|
+
|
|
16
|
+
HELP = """EasyClassifier - machine learning classification without programming.
|
|
17
|
+
|
|
18
|
+
Usage:
|
|
19
|
+
easyclassifier start the step-by-step wizard
|
|
20
|
+
python -m easyclassifier the same, if the 'easyclassifier' command is
|
|
21
|
+
not found
|
|
22
|
+
|
|
23
|
+
Options:
|
|
24
|
+
--version show the installed version
|
|
25
|
+
--help show this help
|
|
26
|
+
|
|
27
|
+
The wizard asks for a data file (.csv or .xlsx). Type 'demo' at that
|
|
28
|
+
question to try it with a built-in example. Results are saved in a new
|
|
29
|
+
folder inside 'Results', in the folder you started from.
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def main(argv=None):
|
|
34
|
+
args = sys.argv[1:] if argv is None else list(argv)
|
|
35
|
+
if any(a in ("-h", "--help", "/?") for a in args):
|
|
36
|
+
print(HELP)
|
|
37
|
+
return 0
|
|
38
|
+
if any(a in ("-V", "--version") for a in args):
|
|
39
|
+
from . import __version__
|
|
40
|
+
print(f"EasyClassifier {__version__} "
|
|
41
|
+
f"(Python {sys.version.split()[0]})")
|
|
42
|
+
return 0
|
|
43
|
+
if args:
|
|
44
|
+
print(f"Unknown option: {' '.join(args)}\n")
|
|
45
|
+
print(HELP)
|
|
46
|
+
return 2
|
|
47
|
+
|
|
48
|
+
import warnings
|
|
49
|
+
# Keep the screen clean for non-programmers; details go to the log.
|
|
50
|
+
warnings.filterwarnings("ignore")
|
|
51
|
+
from . import run
|
|
52
|
+
try:
|
|
53
|
+
run()
|
|
54
|
+
except KeyboardInterrupt:
|
|
55
|
+
print("\n\nExiting EasyClassifier. Goodbye!")
|
|
56
|
+
return 130
|
|
57
|
+
return 0
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
if __name__ == "__main__":
|
|
61
|
+
sys.exit(main())
|
|
@@ -0,0 +1,142 @@
|
|
|
1
|
+
"""Dataset loading and inspection."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
from dataclasses import dataclass, field
|
|
7
|
+
from typing import List, Optional
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import pandas as pd
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@dataclass
|
|
14
|
+
class Inspection:
|
|
15
|
+
"""Summary of a dataset's structure and quality."""
|
|
16
|
+
|
|
17
|
+
n_rows: int
|
|
18
|
+
n_cols: int
|
|
19
|
+
columns: List[str]
|
|
20
|
+
numeric_cols: List[str]
|
|
21
|
+
categorical_cols: List[str]
|
|
22
|
+
missing_total: int
|
|
23
|
+
missing_by_col: dict
|
|
24
|
+
duplicate_rows: int
|
|
25
|
+
columns_with_missing: List[str] = field(default_factory=list)
|
|
26
|
+
|
|
27
|
+
@property
|
|
28
|
+
def has_missing(self) -> bool:
|
|
29
|
+
return self.missing_total > 0
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
def has_duplicates(self) -> bool:
|
|
33
|
+
return self.duplicate_rows > 0
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def clean_path(raw: str) -> str:
|
|
37
|
+
"""Turn what the user typed or dragged into the window into a path.
|
|
38
|
+
|
|
39
|
+
* Dragging a file into a terminal adds quotes (Windows) or escapes
|
|
40
|
+
spaces with a backslash (macOS / Linux); PowerShell may add "& ".
|
|
41
|
+
* ``~`` means the home folder.
|
|
42
|
+
* If the name has no extension, .csv / .xlsx are tried.
|
|
43
|
+
"""
|
|
44
|
+
p = raw.strip()
|
|
45
|
+
if p.startswith("& "):
|
|
46
|
+
p = p[2:].strip()
|
|
47
|
+
if len(p) >= 2 and p[0] == p[-1] and p[0] in "\"'":
|
|
48
|
+
p = p[1:-1]
|
|
49
|
+
if os.sep == "/" and "\\ " in p:
|
|
50
|
+
p = p.replace("\\ ", " ")
|
|
51
|
+
p = os.path.expanduser(p)
|
|
52
|
+
if not os.path.splitext(p)[1] and not os.path.isfile(p):
|
|
53
|
+
for ext in (".csv", ".xlsx"):
|
|
54
|
+
if os.path.isfile(p + ext):
|
|
55
|
+
return p + ext
|
|
56
|
+
return p
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
EXCEL_EXTENSIONS = (".xlsx", ".xlsm")
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def load_csv(path: str) -> pd.DataFrame:
|
|
63
|
+
"""Load a CSV file (or the first sheet of an Excel .xlsx file),
|
|
64
|
+
tolerating common separators and encodings."""
|
|
65
|
+
if not os.path.isfile(path):
|
|
66
|
+
raise FileNotFoundError(path)
|
|
67
|
+
ext = os.path.splitext(path)[1].lower()
|
|
68
|
+
if ext in EXCEL_EXTENSIONS:
|
|
69
|
+
return pd.read_excel(path)
|
|
70
|
+
if ext == ".xls":
|
|
71
|
+
raise ValueError("old .xls Excel files are not supported; please "
|
|
72
|
+
"save the sheet as .xlsx or .csv in Excel "
|
|
73
|
+
"(File > Save As)")
|
|
74
|
+
# Try a few sensible fallbacks for real-world CSVs.
|
|
75
|
+
attempts = [
|
|
76
|
+
dict(),
|
|
77
|
+
dict(sep=";"),
|
|
78
|
+
dict(encoding="latin-1"),
|
|
79
|
+
dict(sep=";", encoding="latin-1"),
|
|
80
|
+
]
|
|
81
|
+
for kwargs in attempts:
|
|
82
|
+
try:
|
|
83
|
+
df = pd.read_csv(path, **kwargs)
|
|
84
|
+
except Exception: # noqa: BLE001
|
|
85
|
+
continue
|
|
86
|
+
if df.shape[1] > 1:
|
|
87
|
+
if kwargs.get("sep") == ";":
|
|
88
|
+
df = fix_decimal_commas(df)
|
|
89
|
+
return df
|
|
90
|
+
# Fall back to the plain read and let its error (if any) surface.
|
|
91
|
+
return pd.read_csv(path)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def fix_decimal_commas(df: pd.DataFrame) -> pd.DataFrame:
|
|
95
|
+
"""Semicolon-separated files (Excel in many countries) write decimals
|
|
96
|
+
with a comma: 3,5. Turn such text columns into numbers when *every*
|
|
97
|
+
value converts; anything else is left unchanged."""
|
|
98
|
+
df = df.copy()
|
|
99
|
+
for c in df.columns:
|
|
100
|
+
if pd.api.types.is_numeric_dtype(df[c]):
|
|
101
|
+
continue
|
|
102
|
+
s = df[c].dropna().astype(str).str.strip()
|
|
103
|
+
if s.empty or not s.str.contains(",", regex=False).any():
|
|
104
|
+
continue
|
|
105
|
+
conv = pd.to_numeric(s.str.replace(",", ".", regex=False),
|
|
106
|
+
errors="coerce")
|
|
107
|
+
if conv.notna().all():
|
|
108
|
+
df[c] = pd.to_numeric(
|
|
109
|
+
df[c].astype(str).str.strip().str.replace(",", ".",
|
|
110
|
+
regex=False),
|
|
111
|
+
errors="coerce")
|
|
112
|
+
return df
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def inspect(df: pd.DataFrame) -> Inspection:
|
|
116
|
+
"""Produce an :class:`Inspection` summary of a dataframe."""
|
|
117
|
+
numeric_cols = df.select_dtypes(include=[np.number]).columns.tolist()
|
|
118
|
+
categorical_cols = [c for c in df.columns if c not in numeric_cols]
|
|
119
|
+
missing_by_col = {
|
|
120
|
+
c: int(df[c].isna().sum()) for c in df.columns if df[c].isna().any()
|
|
121
|
+
}
|
|
122
|
+
return Inspection(
|
|
123
|
+
n_rows=int(df.shape[0]),
|
|
124
|
+
n_cols=int(df.shape[1]),
|
|
125
|
+
columns=df.columns.tolist(),
|
|
126
|
+
numeric_cols=numeric_cols,
|
|
127
|
+
categorical_cols=categorical_cols,
|
|
128
|
+
missing_total=int(df.isna().sum().sum()),
|
|
129
|
+
missing_by_col=missing_by_col,
|
|
130
|
+
duplicate_rows=int(df.duplicated().sum()),
|
|
131
|
+
columns_with_missing=list(missing_by_col.keys()),
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def guess_target(df: pd.DataFrame) -> Optional[str]:
|
|
136
|
+
"""Best guess at which column holds the class labels.
|
|
137
|
+
|
|
138
|
+
Delegates to :func:`easyclassifier.target.suggest_target`, which only
|
|
139
|
+
suggests columns that look like groups (never IDs or measurements).
|
|
140
|
+
"""
|
|
141
|
+
from .target import suggest_target
|
|
142
|
+
return suggest_target(df)
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""Built-in demo dataset so users can try the tool without their own CSV."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def load_demo() -> pd.DataFrame:
|
|
9
|
+
"""Return a small, friendly classification dataset (the Iris flowers)."""
|
|
10
|
+
try:
|
|
11
|
+
from sklearn.datasets import load_iris
|
|
12
|
+
data = load_iris(as_frame=True)
|
|
13
|
+
df = data.frame.copy()
|
|
14
|
+
# Give it human-friendly column names and a text target.
|
|
15
|
+
df = df.rename(columns={
|
|
16
|
+
"sepal length (cm)": "sepal_length",
|
|
17
|
+
"sepal width (cm)": "sepal_width",
|
|
18
|
+
"petal length (cm)": "petal_length",
|
|
19
|
+
"petal width (cm)": "petal_width",
|
|
20
|
+
})
|
|
21
|
+
df["Species"] = pd.Categorical.from_codes(
|
|
22
|
+
data.target, data.target_names
|
|
23
|
+
).astype(str)
|
|
24
|
+
df = df.drop(columns=["target"])
|
|
25
|
+
return df
|
|
26
|
+
except Exception: # noqa: BLE001
|
|
27
|
+
# Extremely small fallback if sklearn datasets are unavailable.
|
|
28
|
+
return pd.DataFrame({
|
|
29
|
+
"x1": [1, 2, 3, 4, 5, 6],
|
|
30
|
+
"x2": [6, 5, 4, 3, 2, 1],
|
|
31
|
+
"Species": ["A", "A", "A", "B", "B", "B"],
|
|
32
|
+
})
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
"""Learning curve and plain-language interpretations of results.
|
|
2
|
+
|
|
3
|
+
* ``learning_curve_data`` - how the score changes as more training rows are
|
|
4
|
+
used ("would more data help?").
|
|
5
|
+
* ``interpret_learning_curve`` / ``interpret_comparison`` - one or two
|
|
6
|
+
sentences a non-specialist can act on.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from typing import Dict, List, Optional
|
|
12
|
+
|
|
13
|
+
import numpy as np
|
|
14
|
+
from sklearn.model_selection import StratifiedKFold, learning_curve
|
|
15
|
+
|
|
16
|
+
from .evaluation import SELECTION_LABEL
|
|
17
|
+
|
|
18
|
+
TRAIN_SIZES = np.linspace(0.2, 1.0, 5)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def learning_curve_data(spec, X, y, seed: int = 42) -> Optional[Dict]:
|
|
22
|
+
"""Balanced accuracy on training and validation rows for 20%...100% of
|
|
23
|
+
the available training rows (5-fold cross-validation). Returns None if
|
|
24
|
+
the data are too small to compute it."""
|
|
25
|
+
y = np.asarray(y)
|
|
26
|
+
smallest = int(np.min(np.unique(y, return_counts=True)[1]))
|
|
27
|
+
n_splits = max(2, min(5, smallest))
|
|
28
|
+
if len(y) < 20 or smallest < 2:
|
|
29
|
+
return None
|
|
30
|
+
cv = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
|
|
31
|
+
sizes, train, val = learning_curve(
|
|
32
|
+
spec.factory(), X, y, train_sizes=TRAIN_SIZES, cv=cv,
|
|
33
|
+
scoring="balanced_accuracy", shuffle=True, random_state=seed,
|
|
34
|
+
error_score=np.nan)
|
|
35
|
+
keep = ~np.all(np.isnan(val), axis=1)
|
|
36
|
+
if keep.sum() < 2:
|
|
37
|
+
return None
|
|
38
|
+
return {
|
|
39
|
+
"sizes": sizes[keep],
|
|
40
|
+
"train_mean": np.nanmean(train[keep], axis=1),
|
|
41
|
+
"train_std": np.nanstd(train[keep], axis=1),
|
|
42
|
+
"val_mean": np.nanmean(val[keep], axis=1),
|
|
43
|
+
"val_std": np.nanstd(val[keep], axis=1),
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def interpret_learning_curve(d: Dict) -> str:
|
|
48
|
+
"""Would more data probably help? Based on the last part of the curve."""
|
|
49
|
+
v = d["val_mean"]
|
|
50
|
+
gain = float(v[-1] - v[-2])
|
|
51
|
+
gap = float(d["train_mean"][-1] - v[-1])
|
|
52
|
+
if gain > 0.01:
|
|
53
|
+
s = ("The score was still rising when all rows were used, so "
|
|
54
|
+
"collecting more data would probably improve the results.")
|
|
55
|
+
elif gain < -0.01:
|
|
56
|
+
s = ("The score did not improve at the end of the curve; more data "
|
|
57
|
+
"of the same kind is unlikely to help much.")
|
|
58
|
+
else:
|
|
59
|
+
s = ("The score had levelled off, so more data of the same kind is "
|
|
60
|
+
"unlikely to improve the results much.")
|
|
61
|
+
if gap > 0.15:
|
|
62
|
+
s += (" The model does much better on its training rows than on new "
|
|
63
|
+
"rows, a sign that it partly memorises the training data.")
|
|
64
|
+
return s
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def interpret_comparison(results: List) -> str:
|
|
68
|
+
"""Is the best classifier clearly better than the runner-up?"""
|
|
69
|
+
if len(results) < 2:
|
|
70
|
+
return ""
|
|
71
|
+
a, b = results[0], results[1]
|
|
72
|
+
diff = a.selection_score - b.selection_score
|
|
73
|
+
spread = (float(np.std(a.fold_scores)) if len(a.fold_scores) > 1
|
|
74
|
+
else None)
|
|
75
|
+
label = SELECTION_LABEL.lower()
|
|
76
|
+
if diff < 0.0005:
|
|
77
|
+
tied = [r.classifier_name for r in results
|
|
78
|
+
if a.selection_score - r.selection_score < 0.0005]
|
|
79
|
+
return (f"{', '.join(tied[:-1])} and {tied[-1]} had the same "
|
|
80
|
+
f"{label}; {a.classifier_name} was selected because it is "
|
|
81
|
+
"listed first. They perform equally well on this data.")
|
|
82
|
+
if spread is None:
|
|
83
|
+
return (f"{a.classifier_name} had the highest {label}; with a "
|
|
84
|
+
"single test split the uncertainty of this ranking cannot be "
|
|
85
|
+
"judged.")
|
|
86
|
+
if diff < spread:
|
|
87
|
+
return (f"{a.classifier_name} had the highest {label}, but its lead "
|
|
88
|
+
f"over {b.classifier_name} ({diff * 100:.1f} percentage "
|
|
89
|
+
f"points) is smaller than the variation between test folds "
|
|
90
|
+
f"({spread * 100:.1f} points). The top classifiers perform "
|
|
91
|
+
"similarly; the choice between them is not clear-cut.")
|
|
92
|
+
return (f"{a.classifier_name} had the highest {label}, ahead of "
|
|
93
|
+
f"{b.classifier_name} by {diff * 100:.1f} percentage points, "
|
|
94
|
+
"which is larger than the variation between test folds "
|
|
95
|
+
f"({spread * 100:.1f} points).")
|
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
"""Distance metrics for the KNN classifier, including the Hassanat distance.
|
|
2
|
+
|
|
3
|
+
Hassanat distance (per dimension i):
|
|
4
|
+
|
|
5
|
+
if min(a_i, b_i) >= 0 (normal form):
|
|
6
|
+
D = 1 - (1 + min) / (1 + max)
|
|
7
|
+
|
|
8
|
+
if min(a_i, b_i) < 0 (signed form, for negative values):
|
|
9
|
+
D = 1 - (1 + min + |min|) / (1 + max + |min|)
|
|
10
|
+
|
|
11
|
+
HD(A, B) = sum_i D(a_i, b_i)
|
|
12
|
+
|
|
13
|
+
Each dimension contributes at most 1, which makes the metric invariant to
|
|
14
|
+
data scale, noise and outliers.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from __future__ import annotations
|
|
18
|
+
|
|
19
|
+
from typing import Dict, List
|
|
20
|
+
|
|
21
|
+
import numpy as np
|
|
22
|
+
from sklearn.base import BaseEstimator, ClassifierMixin
|
|
23
|
+
from sklearn.neighbors import KNeighborsClassifier
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
HASSANAT_CITATION = (
|
|
27
|
+
"Hassanat, A. B. (2014). Dimensionality Invariant Similarity Measure. "
|
|
28
|
+
"Journal of American Science, 10(8). arXiv:1409.0923."
|
|
29
|
+
)
|
|
30
|
+
|
|
31
|
+
HASSANAT_CITATIONS: List[str] = [HASSANAT_CITATION]
|
|
32
|
+
|
|
33
|
+
# key -> (menu label, sklearn metric name or None for custom)
|
|
34
|
+
DISTANCES: Dict[str, tuple] = {
|
|
35
|
+
"hassanat": ("Hassanat distance (robust to outliers)",
|
|
36
|
+
None),
|
|
37
|
+
"euclidean": ("Euclidean distance", "euclidean"),
|
|
38
|
+
"manhattan": ("Manhattan (city-block) distance", "manhattan"),
|
|
39
|
+
"chebyshev": ("Chebyshev distance", "chebyshev"),
|
|
40
|
+
"canberra": ("Canberra distance", "canberra"),
|
|
41
|
+
"cosine": ("Cosine distance", "cosine"),
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
DISTANCE_HELP = {
|
|
45
|
+
"hassanat": "The Hassanat distance compares each feature using the ratio "
|
|
46
|
+
"of its smaller to larger value, so every feature adds at most "
|
|
47
|
+
"1 to the total; one extreme value cannot dominate. "
|
|
48
|
+
"EasyClassifier scales the columns to 0-1 first, which gave "
|
|
49
|
+
"the best results in its benchmarks. Citation: "
|
|
50
|
+
+ HASSANAT_CITATION,
|
|
51
|
+
"euclidean": "Euclidean distance is the straight-line distance between "
|
|
52
|
+
"two points. Sensitive to scale, so scaling is advised.",
|
|
53
|
+
"manhattan": "Manhattan distance adds up the absolute differences of "
|
|
54
|
+
"each feature, like walking city blocks.",
|
|
55
|
+
"chebyshev": "Chebyshev distance uses only the single largest feature "
|
|
56
|
+
"difference.",
|
|
57
|
+
"canberra": "Canberra distance is a weighted version of Manhattan that "
|
|
58
|
+
"is sensitive to small values near zero.",
|
|
59
|
+
"cosine": "Cosine distance compares the direction of two rows, ignoring "
|
|
60
|
+
"their size.",
|
|
61
|
+
}
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
# --------------------------------------------------------------------------- #
|
|
65
|
+
# Hassanat distance
|
|
66
|
+
# --------------------------------------------------------------------------- #
|
|
67
|
+
|
|
68
|
+
def hassanat_distance(a, b) -> float:
|
|
69
|
+
"""Hassanat distance between two 1-D vectors."""
|
|
70
|
+
a = np.asarray(a, dtype=float)
|
|
71
|
+
b = np.asarray(b, dtype=float)
|
|
72
|
+
mn = np.minimum(a, b)
|
|
73
|
+
mx = np.maximum(a, b)
|
|
74
|
+
shift = np.where(mn < 0, -mn, 0.0) # |min| only when min < 0
|
|
75
|
+
return float(np.sum(1.0 - (1.0 + mn + shift) / (1.0 + mx + shift)))
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
MAX_CHUNK_CELLS = 4_000_000 # ~32 MB per temporary array
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def hassanat_matrix(XA, XB, chunk: int = 256) -> np.ndarray:
|
|
82
|
+
"""Pairwise Hassanat distances (rows of XA vs rows of XB), vectorised.
|
|
83
|
+
|
|
84
|
+
Rows of XA are processed in chunks small enough that the temporary
|
|
85
|
+
arrays stay within a few hundred MB, even for large datasets.
|
|
86
|
+
"""
|
|
87
|
+
XA = np.asarray(XA, dtype=float)
|
|
88
|
+
XB = np.asarray(XB, dtype=float)
|
|
89
|
+
per_row = max(1, XB.shape[0] * XB.shape[1])
|
|
90
|
+
chunk = max(1, min(chunk, MAX_CHUNK_CELLS // per_row))
|
|
91
|
+
out = np.empty((XA.shape[0], XB.shape[0]))
|
|
92
|
+
for s in range(0, XA.shape[0], chunk):
|
|
93
|
+
A = XA[s:s + chunk, None, :]
|
|
94
|
+
mn = np.minimum(A, XB[None, :, :])
|
|
95
|
+
mx = np.maximum(A, XB[None, :, :])
|
|
96
|
+
shift = np.where(mn < 0, -mn, 0.0)
|
|
97
|
+
out[s:s + chunk] = np.sum(
|
|
98
|
+
1.0 - (1.0 + mn + shift) / (1.0 + mx + shift), axis=2)
|
|
99
|
+
return out
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def hassanat_form(X) -> dict:
|
|
103
|
+
"""Report whether the normal or signed form of the formula is used."""
|
|
104
|
+
X = np.asarray(X, dtype=float)
|
|
105
|
+
neg_cols = int(np.sum((X < 0).any(axis=0)))
|
|
106
|
+
if neg_cols == 0:
|
|
107
|
+
return {"form": "normal", "negative_features": 0,
|
|
108
|
+
"text": "Normal form used (all feature values are >= 0)."}
|
|
109
|
+
return {"form": "signed", "negative_features": neg_cols,
|
|
110
|
+
"text": (f"Signed form used: {neg_cols} feature(s) contain "
|
|
111
|
+
"negative values, so the negative-value case "
|
|
112
|
+
"1 - (1+min+|min|)/(1+max+|min|) was applied to them.")}
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class HassanatKNN(BaseEstimator, ClassifierMixin):
|
|
116
|
+
"""K-Nearest Neighbours classifier using the Hassanat distance."""
|
|
117
|
+
|
|
118
|
+
def __init__(self, n_neighbors: int = 5):
|
|
119
|
+
self.n_neighbors = n_neighbors
|
|
120
|
+
|
|
121
|
+
def fit(self, X, y):
|
|
122
|
+
self.X_ = np.asarray(X, dtype=float)
|
|
123
|
+
y = np.asarray(y)
|
|
124
|
+
self.classes_, self.y_ = np.unique(y, return_inverse=True)
|
|
125
|
+
return self
|
|
126
|
+
|
|
127
|
+
def _neighbors(self, X):
|
|
128
|
+
D = hassanat_matrix(X, self.X_)
|
|
129
|
+
k = min(self.n_neighbors, self.X_.shape[0])
|
|
130
|
+
idx = np.argpartition(D, k - 1, axis=1)[:, :k]
|
|
131
|
+
return idx
|
|
132
|
+
|
|
133
|
+
def predict_proba(self, X):
|
|
134
|
+
idx = self._neighbors(X)
|
|
135
|
+
votes = self.y_[idx]
|
|
136
|
+
n_cls = len(self.classes_)
|
|
137
|
+
proba = np.zeros((votes.shape[0], n_cls))
|
|
138
|
+
for c in range(n_cls):
|
|
139
|
+
proba[:, c] = (votes == c).sum(axis=1)
|
|
140
|
+
return proba / proba.sum(axis=1, keepdims=True)
|
|
141
|
+
|
|
142
|
+
def predict(self, X):
|
|
143
|
+
return self.classes_[np.argmax(self.predict_proba(X), axis=1)]
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def make_knn(distance: str = "hassanat", k: int = 5):
|
|
147
|
+
"""Factory returning a KNN classifier for the chosen distance."""
|
|
148
|
+
if distance == "hassanat":
|
|
149
|
+
return HassanatKNN(n_neighbors=k)
|
|
150
|
+
metric = DISTANCES[distance][1]
|
|
151
|
+
return KNeighborsClassifier(n_neighbors=k, metric=metric,
|
|
152
|
+
algorithm="brute")
|