geds-python 0.1.0a2__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.
- geds/__init__.py +13 -0
- geds/_backend.py +290 -0
- geds/_estimators.py +385 -0
- geds/py.typed +1 -0
- geds_python-0.1.0a2.dist-info/METADATA +186 -0
- geds_python-0.1.0a2.dist-info/RECORD +8 -0
- geds_python-0.1.0a2.dist-info/WHEEL +4 -0
- geds_python-0.1.0a2.dist-info/licenses/LICENSE +674 -0
geds/__init__.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
"""Python access to the GeDS R package."""
|
|
2
|
+
|
|
3
|
+
from ._backend import BackendUnavailableError, diagnostics
|
|
4
|
+
from ._estimators import GeDSGeneralizedRegressor, GeDSRegressor
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"BackendUnavailableError",
|
|
8
|
+
"GeDSGeneralizedRegressor",
|
|
9
|
+
"GeDSRegressor",
|
|
10
|
+
"diagnostics",
|
|
11
|
+
]
|
|
12
|
+
|
|
13
|
+
__version__ = "0.1.0a2"
|
geds/_backend.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
1
|
+
"""Lazy, narrowly scoped access to the R implementation of GeDS."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from contextlib import contextmanager
|
|
6
|
+
import ctypes
|
|
7
|
+
from importlib.metadata import version
|
|
8
|
+
import os
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
import re
|
|
11
|
+
import shutil
|
|
12
|
+
import subprocess
|
|
13
|
+
import sys
|
|
14
|
+
from typing import Any, Iterator
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
import pandas as pd
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class BackendUnavailableError(RuntimeError):
|
|
21
|
+
"""Raised when R, rpy2, or the GeDS R package is unavailable."""
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
_DLL_HANDLES: list[Any] = []
|
|
25
|
+
MIN_GEDS_VERSION = "0.3.6"
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _version_key(path: Path) -> tuple[int, ...]:
|
|
29
|
+
match = re.search(r"R-(\d+(?:\.\d+)*)$", path.name)
|
|
30
|
+
return tuple(int(part) for part in match.group(1).split(".")) if match else ()
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _discover_r_home() -> Path:
|
|
34
|
+
configured = os.environ.get("R_HOME")
|
|
35
|
+
if configured and (Path(configured) / "bin").is_dir():
|
|
36
|
+
return Path(configured)
|
|
37
|
+
|
|
38
|
+
if os.name == "nt":
|
|
39
|
+
roots = [Path(os.environ.get("ProgramFiles", r"C:\Program Files")) / "R"]
|
|
40
|
+
candidates = [
|
|
41
|
+
path
|
|
42
|
+
for root in roots
|
|
43
|
+
if root.is_dir()
|
|
44
|
+
for path in root.glob("R-*")
|
|
45
|
+
if (path / "bin" / "Rscript.exe").is_file()
|
|
46
|
+
]
|
|
47
|
+
if candidates:
|
|
48
|
+
return max(candidates, key=_version_key)
|
|
49
|
+
|
|
50
|
+
rscript = shutil.which("Rscript")
|
|
51
|
+
if rscript:
|
|
52
|
+
executable = Path(rscript).resolve()
|
|
53
|
+
try:
|
|
54
|
+
reported_home = subprocess.run(
|
|
55
|
+
[str(executable), "--vanilla", "-e", "cat(R.home())"],
|
|
56
|
+
check=True,
|
|
57
|
+
capture_output=True,
|
|
58
|
+
text=True,
|
|
59
|
+
).stdout.strip()
|
|
60
|
+
except (OSError, subprocess.CalledProcessError):
|
|
61
|
+
reported_home = ""
|
|
62
|
+
if reported_home and (Path(reported_home) / "bin").is_dir():
|
|
63
|
+
return Path(reported_home)
|
|
64
|
+
|
|
65
|
+
raise BackendUnavailableError(
|
|
66
|
+
"R was not found. Install R and either put Rscript on PATH or set R_HOME."
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _configure_r_process() -> Path:
|
|
71
|
+
"""Configure the current process before rpy2 imports R's shared library."""
|
|
72
|
+
r_home = _discover_r_home()
|
|
73
|
+
os.environ["R_HOME"] = str(r_home)
|
|
74
|
+
|
|
75
|
+
if os.name == "nt":
|
|
76
|
+
r_bin = r_home / "bin" / "x64"
|
|
77
|
+
if not r_bin.is_dir():
|
|
78
|
+
r_bin = r_home / "bin"
|
|
79
|
+
os.environ["PATH"] = os.pathsep.join(
|
|
80
|
+
(str(r_bin), str(r_home / "bin"), os.environ.get("PATH", ""))
|
|
81
|
+
)
|
|
82
|
+
if hasattr(os, "add_dll_directory"):
|
|
83
|
+
_DLL_HANDLES.append(os.add_dll_directory(str(r_bin)))
|
|
84
|
+
# R loads package DLLs itself. On current Windows/Python combinations,
|
|
85
|
+
# registering the directory with the OS is also required for their
|
|
86
|
+
# transitive dependencies (for example stats.dll -> R.dll).
|
|
87
|
+
if not ctypes.windll.kernel32.SetDllDirectoryW(str(r_bin)):
|
|
88
|
+
raise BackendUnavailableError(
|
|
89
|
+
f"Windows could not register R's DLL directory: {r_bin}"
|
|
90
|
+
)
|
|
91
|
+
# Explicitly retain R's core DLLs as well. R 4.6 can otherwise start
|
|
92
|
+
# successfully while later failing to load recommended packages such
|
|
93
|
+
# as stats because their Rblas/Rlapack dependencies are not found.
|
|
94
|
+
for dll_name in ("R.dll", "Rblas.dll", "Rlapack.dll", "Riconv.dll"):
|
|
95
|
+
dll_path = r_bin / dll_name
|
|
96
|
+
if dll_path.is_file():
|
|
97
|
+
try:
|
|
98
|
+
_DLL_HANDLES.append(ctypes.WinDLL(str(dll_path)))
|
|
99
|
+
except OSError as exc:
|
|
100
|
+
raise BackendUnavailableError(
|
|
101
|
+
f"Windows could not load R's core library: {dll_path}"
|
|
102
|
+
) from exc
|
|
103
|
+
|
|
104
|
+
return r_home
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class RBackend:
|
|
108
|
+
"""Own the embedded-R session and keep fitted models opaque."""
|
|
109
|
+
|
|
110
|
+
def __init__(self) -> None:
|
|
111
|
+
self.r_home = _configure_r_process()
|
|
112
|
+
rpy2_situation = None
|
|
113
|
+
original_get_r_flags = None
|
|
114
|
+
try:
|
|
115
|
+
# rpy2 3.6.x probes ``R CMD config --ldflags`` before loading
|
|
116
|
+
# R.dll. Some Windows R builds return no stdout for that query;
|
|
117
|
+
# rpy2 then raises IndexError instead of taking its normal
|
|
118
|
+
# Windows DLL-directory fallback. Translate only that empty-output
|
|
119
|
+
# case into the exception the fallback already handles.
|
|
120
|
+
if os.name == "nt":
|
|
121
|
+
import rpy2.situation as rpy2_situation
|
|
122
|
+
|
|
123
|
+
original_get_r_flags = rpy2_situation.get_r_flags
|
|
124
|
+
|
|
125
|
+
def get_r_flags_with_windows_fallback(*args: Any, **kwargs: Any) -> Any:
|
|
126
|
+
try:
|
|
127
|
+
return original_get_r_flags(*args, **kwargs)
|
|
128
|
+
except IndexError as exc:
|
|
129
|
+
raise subprocess.CalledProcessError(1, "R CMD config") from exc
|
|
130
|
+
|
|
131
|
+
rpy2_situation.get_r_flags = get_r_flags_with_windows_fallback
|
|
132
|
+
|
|
133
|
+
import rpy2.robjects as ro
|
|
134
|
+
from rpy2.rinterface_lib import openrlib
|
|
135
|
+
from rpy2.robjects import default_converter, pandas2ri
|
|
136
|
+
from rpy2.robjects.packages import importr
|
|
137
|
+
except Exception as exc: # pragma: no cover - depends on host setup
|
|
138
|
+
raise BackendUnavailableError(
|
|
139
|
+
"rpy2 could not initialize embedded R. Run geds.diagnostics() "
|
|
140
|
+
"and verify that R and rpy2 use compatible versions."
|
|
141
|
+
) from exc
|
|
142
|
+
finally:
|
|
143
|
+
if rpy2_situation is not None and original_get_r_flags is not None:
|
|
144
|
+
rpy2_situation.get_r_flags = original_get_r_flags
|
|
145
|
+
|
|
146
|
+
self.ro = ro
|
|
147
|
+
self._openrlib = openrlib
|
|
148
|
+
self._converter = default_converter + pandas2ri.converter
|
|
149
|
+
|
|
150
|
+
extra_library = os.environ.get("GEDS_R_LIBRARY")
|
|
151
|
+
if extra_library:
|
|
152
|
+
current = list(ro.r(".libPaths()"))
|
|
153
|
+
ro.r[".libPaths"](ro.StrVector([extra_library, *current]))
|
|
154
|
+
|
|
155
|
+
try:
|
|
156
|
+
self.geds = importr("GeDS")
|
|
157
|
+
self.stats = importr("stats")
|
|
158
|
+
except Exception as exc:
|
|
159
|
+
raise BackendUnavailableError(
|
|
160
|
+
"The GeDS R package could not be loaded. Install GeDS into a "
|
|
161
|
+
"library visible to this R installation."
|
|
162
|
+
) from exc
|
|
163
|
+
|
|
164
|
+
self.geds_version = str(
|
|
165
|
+
self.ro.r("as.character(packageVersion('GeDS'))")[0]
|
|
166
|
+
)
|
|
167
|
+
is_supported = bool(
|
|
168
|
+
self.ro.r(
|
|
169
|
+
"packageVersion('GeDS') >= "
|
|
170
|
+
f"package_version('{MIN_GEDS_VERSION}')"
|
|
171
|
+
)[0]
|
|
172
|
+
)
|
|
173
|
+
if not is_supported:
|
|
174
|
+
raise BackendUnavailableError(
|
|
175
|
+
f"GeDS {self.geds_version} is installed, but geds-python "
|
|
176
|
+
f"requires GeDS >= {MIN_GEDS_VERSION}."
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
@contextmanager
|
|
180
|
+
def locked(self) -> Iterator[None]:
|
|
181
|
+
"""Serialize access to R's process-global runtime."""
|
|
182
|
+
with self._openrlib.rlock:
|
|
183
|
+
yield
|
|
184
|
+
|
|
185
|
+
def dataframe_to_r(self, frame: pd.DataFrame) -> Any:
|
|
186
|
+
with self.locked(), self._converter.context():
|
|
187
|
+
return self.ro.conversion.get_conversion().py2rpy(frame)
|
|
188
|
+
|
|
189
|
+
def vector(self, values: Any) -> Any:
|
|
190
|
+
return self.ro.FloatVector(np.asarray(values, dtype=float))
|
|
191
|
+
|
|
192
|
+
def formula(self, expression: str) -> Any:
|
|
193
|
+
return self.ro.Formula(expression)
|
|
194
|
+
|
|
195
|
+
def family(self, name: str, link: str | None) -> Any:
|
|
196
|
+
normalized = name.lower()
|
|
197
|
+
functions = {
|
|
198
|
+
"gaussian": "gaussian",
|
|
199
|
+
"poisson": "poisson",
|
|
200
|
+
"quasipoisson": "quasipoisson",
|
|
201
|
+
"binomial": "binomial",
|
|
202
|
+
"quasibinomial": "quasibinomial",
|
|
203
|
+
"gamma": "Gamma",
|
|
204
|
+
}
|
|
205
|
+
if normalized not in functions:
|
|
206
|
+
choices = ", ".join(sorted(functions))
|
|
207
|
+
raise ValueError(f"Unsupported family {name!r}; choose one of: {choices}.")
|
|
208
|
+
function = getattr(self.stats, functions[normalized])
|
|
209
|
+
return function() if link is None else function(link=link)
|
|
210
|
+
|
|
211
|
+
def fit(self, function: str, formula: Any, data: Any, **kwargs: Any) -> Any:
|
|
212
|
+
with self.locked():
|
|
213
|
+
return getattr(self.geds, function)(formula, data=data, **kwargs)
|
|
214
|
+
|
|
215
|
+
def predict(
|
|
216
|
+
self, model: Any, data: Any, order: int, prediction_type: str
|
|
217
|
+
) -> np.ndarray:
|
|
218
|
+
with self.locked():
|
|
219
|
+
result = self.stats.predict(
|
|
220
|
+
model, newdata=data, n=order, type=prediction_type
|
|
221
|
+
)
|
|
222
|
+
return np.asarray(result, dtype=float)
|
|
223
|
+
|
|
224
|
+
def coefficients(self, model: Any, order: int) -> Any:
|
|
225
|
+
with self.locked():
|
|
226
|
+
return self.to_python(self.stats.coef(model, n=order))
|
|
227
|
+
|
|
228
|
+
def knots(self, model: Any, order: int) -> Any:
|
|
229
|
+
with self.locked():
|
|
230
|
+
return self.to_python(
|
|
231
|
+
self.stats.knots(model, n=order, options="internal")
|
|
232
|
+
)
|
|
233
|
+
|
|
234
|
+
def deviance(self, model: Any, order: int) -> float:
|
|
235
|
+
with self.locked():
|
|
236
|
+
return float(self.stats.deviance(model, n=order)[0])
|
|
237
|
+
|
|
238
|
+
def component(self, model: Any, name: str) -> Any:
|
|
239
|
+
try:
|
|
240
|
+
return self.to_python(model.rx2(name))
|
|
241
|
+
except Exception:
|
|
242
|
+
return None
|
|
243
|
+
|
|
244
|
+
def to_python(self, value: Any) -> Any:
|
|
245
|
+
"""Convert extracted values without recursively converting a model."""
|
|
246
|
+
vectors = self.ro.vectors
|
|
247
|
+
if value is self.ro.NULL:
|
|
248
|
+
return None
|
|
249
|
+
if isinstance(value, vectors.ListVector):
|
|
250
|
+
names = list(value.names) if value.names is not self.ro.NULL else []
|
|
251
|
+
converted = [self.to_python(item) for item in value]
|
|
252
|
+
if names and len(names) == len(converted) and all(names):
|
|
253
|
+
return dict(zip(names, converted))
|
|
254
|
+
return converted
|
|
255
|
+
if isinstance(value, vectors.StrVector):
|
|
256
|
+
result = np.asarray(value, dtype=str)
|
|
257
|
+
elif isinstance(
|
|
258
|
+
value, (vectors.FloatVector, vectors.IntVector, vectors.BoolVector)
|
|
259
|
+
):
|
|
260
|
+
result = np.asarray(value)
|
|
261
|
+
else:
|
|
262
|
+
return value
|
|
263
|
+
return result.item() if result.ndim == 0 else result
|
|
264
|
+
|
|
265
|
+
def information(self) -> dict[str, str]:
|
|
266
|
+
with self.locked():
|
|
267
|
+
return {
|
|
268
|
+
"r_home": str(self.r_home),
|
|
269
|
+
"r_version": str(self.ro.r("R.version.string")[0]),
|
|
270
|
+
"geds_version": self.geds_version,
|
|
271
|
+
"minimum_geds_version": MIN_GEDS_VERSION,
|
|
272
|
+
"geds_library": str(self.ro.r("find.package('GeDS')")[0]),
|
|
273
|
+
"rpy2_version": version("rpy2"),
|
|
274
|
+
"python": sys.version.split()[0],
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
_BACKEND: RBackend | None = None
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def get_backend() -> RBackend:
|
|
282
|
+
global _BACKEND
|
|
283
|
+
if _BACKEND is None:
|
|
284
|
+
_BACKEND = RBackend()
|
|
285
|
+
return _BACKEND
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
def diagnostics() -> dict[str, str]:
|
|
289
|
+
"""Return the versions and locations used by the Python/R bridge."""
|
|
290
|
+
return get_backend().information()
|
geds/_estimators.py
ADDED
|
@@ -0,0 +1,385 @@
|
|
|
1
|
+
"""Scikit-learn-style estimators that delegate all statistics to R GeDS."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
import pickle
|
|
7
|
+
from typing import Any, Sequence
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import pandas as pd
|
|
11
|
+
from pandas.api.types import is_numeric_dtype
|
|
12
|
+
from sklearn.base import BaseEstimator, RegressorMixin
|
|
13
|
+
from sklearn.utils.validation import check_is_fitted
|
|
14
|
+
|
|
15
|
+
from ._backend import get_backend
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
FeatureSelector = Sequence[str | int] | None
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class _GeDSBase(RegressorMixin, BaseEstimator):
|
|
22
|
+
_fit_function = ""
|
|
23
|
+
|
|
24
|
+
def _validate_configuration(self) -> None:
|
|
25
|
+
if self.order not in (2, 3, 4):
|
|
26
|
+
raise ValueError("order must be 2, 3, or 4.")
|
|
27
|
+
if not self.higher_order and self.order != 2:
|
|
28
|
+
raise ValueError("order must be 2 when higher_order=False.")
|
|
29
|
+
if self.beta is not None and not 0 <= self.beta <= 1:
|
|
30
|
+
raise ValueError("beta must lie in [0, 1].")
|
|
31
|
+
if not 0 <= self.phi <= 1:
|
|
32
|
+
raise ValueError("phi must lie in [0, 1].")
|
|
33
|
+
if not isinstance(self.q, (int, np.integer)) or self.q < 1:
|
|
34
|
+
raise ValueError("q must be a positive integer.")
|
|
35
|
+
if self.stop_type not in {"SR", "RD", "LR"}:
|
|
36
|
+
raise ValueError("stop_type must be 'SR', 'RD', or 'LR'.")
|
|
37
|
+
for name in ("min_internal_knots", "max_internal_knots"):
|
|
38
|
+
value = getattr(self, name)
|
|
39
|
+
if value is not None and (
|
|
40
|
+
not isinstance(value, (int, np.integer)) or value < 0
|
|
41
|
+
):
|
|
42
|
+
raise ValueError(f"{name} must be a non-negative integer or None.")
|
|
43
|
+
if (
|
|
44
|
+
self.min_internal_knots is not None
|
|
45
|
+
and self.max_internal_knots is not None
|
|
46
|
+
and self.min_internal_knots > self.max_internal_knots
|
|
47
|
+
):
|
|
48
|
+
raise ValueError(
|
|
49
|
+
"min_internal_knots must not exceed max_internal_knots."
|
|
50
|
+
)
|
|
51
|
+
self._validate_range("x_range", self.x_range)
|
|
52
|
+
self._validate_range("y_range", self.y_range)
|
|
53
|
+
|
|
54
|
+
@staticmethod
|
|
55
|
+
def _validate_range(name: str, value: Sequence[float] | None) -> None:
|
|
56
|
+
if value is None:
|
|
57
|
+
return
|
|
58
|
+
array = np.asarray(value, dtype=float)
|
|
59
|
+
if array.shape != (2,) or not np.isfinite(array).all() or array[0] >= array[1]:
|
|
60
|
+
raise ValueError(
|
|
61
|
+
f"{name} must contain two finite, strictly increasing values."
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
@staticmethod
|
|
65
|
+
def _frame(X: Any) -> tuple[pd.DataFrame, bool]:
|
|
66
|
+
if isinstance(X, pd.DataFrame):
|
|
67
|
+
if X.columns.has_duplicates:
|
|
68
|
+
raise ValueError("X must not contain duplicate column names.")
|
|
69
|
+
return X.copy(), True
|
|
70
|
+
array = np.asarray(X)
|
|
71
|
+
if array.ndim != 2:
|
|
72
|
+
raise ValueError("X must be a two-dimensional array or DataFrame.")
|
|
73
|
+
return (
|
|
74
|
+
pd.DataFrame(array, columns=[f"x{i}" for i in range(array.shape[1])]),
|
|
75
|
+
False,
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
@staticmethod
|
|
79
|
+
def _resolve(
|
|
80
|
+
selectors: FeatureSelector, columns: list[Any], *, default: list[int]
|
|
81
|
+
) -> list[int]:
|
|
82
|
+
if selectors is None:
|
|
83
|
+
return default
|
|
84
|
+
if isinstance(selectors, (str, int)):
|
|
85
|
+
selectors = [selectors]
|
|
86
|
+
indices: list[int] = []
|
|
87
|
+
for selector in selectors:
|
|
88
|
+
if isinstance(selector, str):
|
|
89
|
+
if selector not in columns:
|
|
90
|
+
raise ValueError(f"Unknown feature {selector!r}.")
|
|
91
|
+
index = columns.index(selector)
|
|
92
|
+
elif isinstance(selector, (int, np.integer)):
|
|
93
|
+
index = int(selector)
|
|
94
|
+
if index < 0 or index >= len(columns):
|
|
95
|
+
raise ValueError(f"Feature index {index} is out of range.")
|
|
96
|
+
else:
|
|
97
|
+
raise TypeError("Feature selectors must be column names or indices.")
|
|
98
|
+
if index not in indices:
|
|
99
|
+
indices.append(index)
|
|
100
|
+
return indices
|
|
101
|
+
|
|
102
|
+
def _prepare_training_data(
|
|
103
|
+
self, X: Any, y: Any, sample_weight: Any
|
|
104
|
+
) -> tuple[pd.DataFrame, str, Any | None]:
|
|
105
|
+
frame, named_input = self._frame(X)
|
|
106
|
+
y_array = np.asarray(y, dtype=float)
|
|
107
|
+
if y_array.ndim != 1 or len(y_array) != len(frame):
|
|
108
|
+
raise ValueError("y must be one-dimensional and have the same length as X.")
|
|
109
|
+
if not np.isfinite(y_array).all():
|
|
110
|
+
raise ValueError("y must contain only finite values.")
|
|
111
|
+
|
|
112
|
+
columns = list(frame.columns)
|
|
113
|
+
spline = self._resolve(
|
|
114
|
+
self.spline_features, columns, default=list(range(len(columns)))
|
|
115
|
+
)
|
|
116
|
+
if not spline:
|
|
117
|
+
raise ValueError("At least one spline feature is required.")
|
|
118
|
+
remainder = [index for index in range(len(columns)) if index not in spline]
|
|
119
|
+
linear = self._resolve(self.linear_features, columns, default=remainder)
|
|
120
|
+
overlap = sorted(set(spline).intersection(linear))
|
|
121
|
+
if overlap:
|
|
122
|
+
raise ValueError("Spline and linear feature selections must not overlap.")
|
|
123
|
+
non_numeric = [
|
|
124
|
+
columns[index]
|
|
125
|
+
for index in spline
|
|
126
|
+
if not is_numeric_dtype(frame.iloc[:, index])
|
|
127
|
+
]
|
|
128
|
+
if non_numeric:
|
|
129
|
+
raise TypeError(f"Spline features must be numeric; got {non_numeric}.")
|
|
130
|
+
if frame.isna().to_numpy().any():
|
|
131
|
+
raise ValueError("X must not contain missing values.")
|
|
132
|
+
|
|
133
|
+
self.n_features_in_ = frame.shape[1]
|
|
134
|
+
self._input_columns_ = columns
|
|
135
|
+
self._named_input_ = named_input
|
|
136
|
+
if named_input and all(isinstance(column, str) for column in columns):
|
|
137
|
+
self.feature_names_in_ = np.asarray(columns, dtype=object)
|
|
138
|
+
self._internal_columns_ = [f"x{index}" for index in range(len(columns))]
|
|
139
|
+
self._spline_indices_ = spline
|
|
140
|
+
self._linear_indices_ = linear
|
|
141
|
+
self.spline_features_ = np.asarray(
|
|
142
|
+
[columns[index] for index in spline], dtype=object
|
|
143
|
+
)
|
|
144
|
+
self.linear_features_ = np.asarray(
|
|
145
|
+
[columns[index] for index in linear], dtype=object
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
internal = frame.copy()
|
|
149
|
+
internal.columns = self._internal_columns_
|
|
150
|
+
internal.insert(0, "response", y_array)
|
|
151
|
+
spline_term = "f(" + ", ".join(self._internal_columns_[i] for i in spline) + ")"
|
|
152
|
+
linear_terms = [self._internal_columns_[i] for i in linear]
|
|
153
|
+
formula = "response ~ " + " + ".join([spline_term, *linear_terms])
|
|
154
|
+
self.formula_ = formula
|
|
155
|
+
|
|
156
|
+
weights = None
|
|
157
|
+
if sample_weight is not None:
|
|
158
|
+
weight_array = np.asarray(sample_weight, dtype=float)
|
|
159
|
+
if weight_array.ndim != 1 or len(weight_array) != len(frame):
|
|
160
|
+
raise ValueError(
|
|
161
|
+
"sample_weight must be one-dimensional and match the length of X."
|
|
162
|
+
)
|
|
163
|
+
if not np.isfinite(weight_array).all() or np.any(weight_array < 0):
|
|
164
|
+
raise ValueError(
|
|
165
|
+
"sample_weight must contain only finite, non-negative values."
|
|
166
|
+
)
|
|
167
|
+
weights = get_backend().vector(weight_array)
|
|
168
|
+
return internal, formula, weights
|
|
169
|
+
|
|
170
|
+
def _prepare_new_data(self, X: Any) -> pd.DataFrame:
|
|
171
|
+
check_is_fitted(self, "_r_model_")
|
|
172
|
+
frame, named_input = self._frame(X)
|
|
173
|
+
if frame.shape[1] != self.n_features_in_:
|
|
174
|
+
raise ValueError(
|
|
175
|
+
f"X has {frame.shape[1]} features; expected {self.n_features_in_}."
|
|
176
|
+
)
|
|
177
|
+
if self._named_input_:
|
|
178
|
+
missing = [
|
|
179
|
+
name for name in self._input_columns_ if name not in frame.columns
|
|
180
|
+
]
|
|
181
|
+
if missing:
|
|
182
|
+
raise ValueError(f"X is missing columns: {missing}.")
|
|
183
|
+
frame = frame.loc[:, self._input_columns_]
|
|
184
|
+
elif named_input:
|
|
185
|
+
frame = frame.iloc[:, : self.n_features_in_]
|
|
186
|
+
if frame.isna().to_numpy().any():
|
|
187
|
+
raise ValueError("X must not contain missing values.")
|
|
188
|
+
non_numeric = [
|
|
189
|
+
self._input_columns_[index]
|
|
190
|
+
for index in self._spline_indices_
|
|
191
|
+
if not is_numeric_dtype(frame.iloc[:, index])
|
|
192
|
+
]
|
|
193
|
+
if non_numeric:
|
|
194
|
+
raise TypeError(f"Spline features must be numeric; got {non_numeric}.")
|
|
195
|
+
frame.columns = self._internal_columns_
|
|
196
|
+
return frame
|
|
197
|
+
|
|
198
|
+
def _common_fit_kwargs(self, weights: Any | None) -> dict[str, Any]:
|
|
199
|
+
kwargs: dict[str, Any] = {
|
|
200
|
+
"phi": self.phi,
|
|
201
|
+
"q": self.q,
|
|
202
|
+
"show_iters": self.verbose,
|
|
203
|
+
"stoptype": self.stop_type,
|
|
204
|
+
"higher_order": self.higher_order,
|
|
205
|
+
}
|
|
206
|
+
optional = {
|
|
207
|
+
"beta": self.beta,
|
|
208
|
+
"min_intknots": self.min_internal_knots,
|
|
209
|
+
"max_intknots": self.max_internal_knots,
|
|
210
|
+
"Xextr": (
|
|
211
|
+
get_backend().vector(self.x_range) if self.x_range is not None else None
|
|
212
|
+
),
|
|
213
|
+
"Yextr": (
|
|
214
|
+
get_backend().vector(self.y_range) if self.y_range is not None else None
|
|
215
|
+
),
|
|
216
|
+
"weights": weights,
|
|
217
|
+
}
|
|
218
|
+
kwargs.update(
|
|
219
|
+
{key: value for key, value in optional.items() if value is not None}
|
|
220
|
+
)
|
|
221
|
+
return kwargs
|
|
222
|
+
|
|
223
|
+
def _finish_fit(self, model: Any) -> None:
|
|
224
|
+
backend = get_backend()
|
|
225
|
+
self._r_model_ = model
|
|
226
|
+
self.coef_ = backend.coefficients(model, self.order)
|
|
227
|
+
self.knots_ = backend.knots(model, self.order)
|
|
228
|
+
self.deviance_ = backend.deviance(model, self.order)
|
|
229
|
+
self.n_iter_ = backend.component(model, "iters")
|
|
230
|
+
|
|
231
|
+
def predict(self, X: Any) -> np.ndarray:
|
|
232
|
+
frame = self._prepare_new_data(X)
|
|
233
|
+
backend = get_backend()
|
|
234
|
+
return backend.predict(
|
|
235
|
+
self._r_model_, backend.dataframe_to_r(frame), self.order, "response"
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
def predict_link(self, X: Any) -> np.ndarray:
|
|
239
|
+
"""Return predictions on the link scale."""
|
|
240
|
+
frame = self._prepare_new_data(X)
|
|
241
|
+
backend = get_backend()
|
|
242
|
+
return backend.predict(
|
|
243
|
+
self._r_model_, backend.dataframe_to_r(frame), self.order, "link"
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
def get_coefficients(self, order: int | None = None) -> Any:
|
|
247
|
+
check_is_fitted(self, "_r_model_")
|
|
248
|
+
selected_order = self.order if order is None else order
|
|
249
|
+
self._validate_requested_order(selected_order)
|
|
250
|
+
return get_backend().coefficients(self._r_model_, selected_order)
|
|
251
|
+
|
|
252
|
+
def get_knots(self, order: int | None = None) -> Any:
|
|
253
|
+
check_is_fitted(self, "_r_model_")
|
|
254
|
+
selected_order = self.order if order is None else order
|
|
255
|
+
self._validate_requested_order(selected_order)
|
|
256
|
+
return get_backend().knots(self._r_model_, selected_order)
|
|
257
|
+
|
|
258
|
+
def _validate_requested_order(self, order: int) -> None:
|
|
259
|
+
if order not in (2, 3, 4):
|
|
260
|
+
raise ValueError("order must be 2, 3, or 4.")
|
|
261
|
+
if not self.higher_order and order != 2:
|
|
262
|
+
raise ValueError("Only order 2 is available when higher_order=False.")
|
|
263
|
+
|
|
264
|
+
def save(self, path: str | Path) -> None:
|
|
265
|
+
"""Serialize the estimator and its opaque R model with pickle."""
|
|
266
|
+
check_is_fitted(self, "_r_model_")
|
|
267
|
+
with Path(path).open("wb") as stream:
|
|
268
|
+
pickle.dump(self, stream, protocol=pickle.HIGHEST_PROTOCOL)
|
|
269
|
+
|
|
270
|
+
@classmethod
|
|
271
|
+
def load(cls, path: str | Path) -> "_GeDSBase":
|
|
272
|
+
"""Load an estimator saved by :meth:`save` from a trusted source."""
|
|
273
|
+
with Path(path).open("rb") as stream:
|
|
274
|
+
model = pickle.load(stream)
|
|
275
|
+
if not isinstance(model, cls):
|
|
276
|
+
raise TypeError(f"The file does not contain a {cls.__name__} estimator.")
|
|
277
|
+
return model
|
|
278
|
+
|
|
279
|
+
|
|
280
|
+
class GeDSRegressor(_GeDSBase):
|
|
281
|
+
"""Normal-response GeDS estimator backed by ``GeDS::NGeDS``."""
|
|
282
|
+
|
|
283
|
+
_fit_function = "NGeDS"
|
|
284
|
+
|
|
285
|
+
def __init__(
|
|
286
|
+
self,
|
|
287
|
+
*,
|
|
288
|
+
spline_features: FeatureSelector = None,
|
|
289
|
+
linear_features: FeatureSelector = None,
|
|
290
|
+
order: int = 3,
|
|
291
|
+
beta: float = 0.5,
|
|
292
|
+
phi: float = 0.99,
|
|
293
|
+
min_internal_knots: int | None = None,
|
|
294
|
+
max_internal_knots: int | None = None,
|
|
295
|
+
q: int = 2,
|
|
296
|
+
x_range: Sequence[float] | None = None,
|
|
297
|
+
y_range: Sequence[float] | None = None,
|
|
298
|
+
stop_type: str = "RD",
|
|
299
|
+
higher_order: bool = True,
|
|
300
|
+
verbose: bool = False,
|
|
301
|
+
) -> None:
|
|
302
|
+
self.spline_features = spline_features
|
|
303
|
+
self.linear_features = linear_features
|
|
304
|
+
self.order = order
|
|
305
|
+
self.beta = beta
|
|
306
|
+
self.phi = phi
|
|
307
|
+
self.min_internal_knots = min_internal_knots
|
|
308
|
+
self.max_internal_knots = max_internal_knots
|
|
309
|
+
self.q = q
|
|
310
|
+
self.x_range = x_range
|
|
311
|
+
self.y_range = y_range
|
|
312
|
+
self.stop_type = stop_type
|
|
313
|
+
self.higher_order = higher_order
|
|
314
|
+
self.verbose = verbose
|
|
315
|
+
|
|
316
|
+
def fit(self, X: Any, y: Any, sample_weight: Any = None) -> "GeDSRegressor":
|
|
317
|
+
self._validate_configuration()
|
|
318
|
+
frame, formula, weights = self._prepare_training_data(X, y, sample_weight)
|
|
319
|
+
backend = get_backend()
|
|
320
|
+
model = backend.fit(
|
|
321
|
+
self._fit_function,
|
|
322
|
+
backend.formula(formula),
|
|
323
|
+
backend.dataframe_to_r(frame),
|
|
324
|
+
**self._common_fit_kwargs(weights),
|
|
325
|
+
)
|
|
326
|
+
self._finish_fit(model)
|
|
327
|
+
return self
|
|
328
|
+
|
|
329
|
+
|
|
330
|
+
class GeDSGeneralizedRegressor(_GeDSBase):
|
|
331
|
+
"""Exponential-family GeDS estimator backed by ``GeDS::GGeDS``."""
|
|
332
|
+
|
|
333
|
+
_fit_function = "GGeDS"
|
|
334
|
+
|
|
335
|
+
def __init__(
|
|
336
|
+
self,
|
|
337
|
+
*,
|
|
338
|
+
family: str = "gaussian",
|
|
339
|
+
link: str | None = None,
|
|
340
|
+
spline_features: FeatureSelector = None,
|
|
341
|
+
linear_features: FeatureSelector = None,
|
|
342
|
+
order: int = 3,
|
|
343
|
+
beta: float | None = None,
|
|
344
|
+
phi: float = 0.99,
|
|
345
|
+
min_internal_knots: int | None = None,
|
|
346
|
+
max_internal_knots: int | None = None,
|
|
347
|
+
q: int = 2,
|
|
348
|
+
x_range: Sequence[float] | None = None,
|
|
349
|
+
y_range: Sequence[float] | None = None,
|
|
350
|
+
stop_type: str = "SR",
|
|
351
|
+
higher_order: bool = True,
|
|
352
|
+
verbose: bool = False,
|
|
353
|
+
) -> None:
|
|
354
|
+
self.family = family
|
|
355
|
+
self.link = link
|
|
356
|
+
self.spline_features = spline_features
|
|
357
|
+
self.linear_features = linear_features
|
|
358
|
+
self.order = order
|
|
359
|
+
self.beta = beta
|
|
360
|
+
self.phi = phi
|
|
361
|
+
self.min_internal_knots = min_internal_knots
|
|
362
|
+
self.max_internal_knots = max_internal_knots
|
|
363
|
+
self.q = q
|
|
364
|
+
self.x_range = x_range
|
|
365
|
+
self.y_range = y_range
|
|
366
|
+
self.stop_type = stop_type
|
|
367
|
+
self.higher_order = higher_order
|
|
368
|
+
self.verbose = verbose
|
|
369
|
+
|
|
370
|
+
def fit(
|
|
371
|
+
self, X: Any, y: Any, sample_weight: Any = None
|
|
372
|
+
) -> "GeDSGeneralizedRegressor":
|
|
373
|
+
self._validate_configuration()
|
|
374
|
+
frame, formula, weights = self._prepare_training_data(X, y, sample_weight)
|
|
375
|
+
backend = get_backend()
|
|
376
|
+
kwargs = self._common_fit_kwargs(weights)
|
|
377
|
+
kwargs["family"] = backend.family(self.family, self.link)
|
|
378
|
+
model = backend.fit(
|
|
379
|
+
self._fit_function,
|
|
380
|
+
backend.formula(formula),
|
|
381
|
+
backend.dataframe_to_r(frame),
|
|
382
|
+
**kwargs,
|
|
383
|
+
)
|
|
384
|
+
self._finish_fit(model)
|
|
385
|
+
return self
|
geds/py.typed
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|