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 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
+