geds-python 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.
geds/__init__.py ADDED
@@ -0,0 +1,25 @@
1
+ """Python access to the GeDS R package."""
2
+
3
+ from ._backend import BackendUnavailableError, diagnostics
4
+ from ._estimators import (
5
+ GeDSBoostRegressor,
6
+ GeDSGAMRegressor,
7
+ GeDSGeneralizedRegressor,
8
+ GeDSRegressor,
9
+ )
10
+ from ._plotting import plot_fit
11
+ from ._validation import GeDSCrossValidationResult, cross_validate_geds
12
+
13
+ __all__ = [
14
+ "BackendUnavailableError",
15
+ "GeDSBoostRegressor",
16
+ "GeDSCrossValidationResult",
17
+ "GeDSGAMRegressor",
18
+ "GeDSGeneralizedRegressor",
19
+ "GeDSRegressor",
20
+ "diagnostics",
21
+ "cross_validate_geds",
22
+ "plot_fit",
23
+ ]
24
+
25
+ __version__ = "0.1.0"
geds/_backend.py ADDED
@@ -0,0 +1,458 @@
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
+ REQUIRED_GEDS_CAPABILITIES = {"offset_glm_v1", "terms_matrix_v1"}
27
+
28
+
29
+ def _version_key(path: Path) -> tuple[int, ...]:
30
+ match = re.search(r"R-(\d+(?:\.\d+)*)$", path.name)
31
+ return tuple(int(part) for part in match.group(1).split(".")) if match else ()
32
+
33
+
34
+ def _discover_r_home() -> Path:
35
+ configured = os.environ.get("R_HOME")
36
+ if configured and (Path(configured) / "bin").is_dir():
37
+ return Path(configured)
38
+
39
+ if os.name == "nt":
40
+ roots = [Path(os.environ.get("ProgramFiles", r"C:\Program Files")) / "R"]
41
+ candidates = [
42
+ path
43
+ for root in roots
44
+ if root.is_dir()
45
+ for path in root.glob("R-*")
46
+ if (path / "bin" / "Rscript.exe").is_file()
47
+ ]
48
+ if candidates:
49
+ return max(candidates, key=_version_key)
50
+
51
+ rscript = shutil.which("Rscript")
52
+ if rscript:
53
+ executable = Path(rscript).resolve()
54
+ try:
55
+ reported_home = subprocess.run(
56
+ [str(executable), "--vanilla", "-e", "cat(R.home())"],
57
+ check=True,
58
+ capture_output=True,
59
+ text=True,
60
+ ).stdout.strip()
61
+ except (OSError, subprocess.CalledProcessError):
62
+ reported_home = ""
63
+ if reported_home and (Path(reported_home) / "bin").is_dir():
64
+ return Path(reported_home)
65
+
66
+ raise BackendUnavailableError(
67
+ "R was not found. Install R and either put Rscript on PATH or set R_HOME."
68
+ )
69
+
70
+
71
+ def _configure_r_process() -> Path:
72
+ """Configure the current process before rpy2 imports R's shared library."""
73
+ r_home = _discover_r_home()
74
+ os.environ["R_HOME"] = str(r_home)
75
+
76
+ if os.name == "nt":
77
+ r_bin = r_home / "bin" / "x64"
78
+ if not r_bin.is_dir():
79
+ r_bin = r_home / "bin"
80
+ os.environ["PATH"] = os.pathsep.join(
81
+ (str(r_bin), str(r_home / "bin"), os.environ.get("PATH", ""))
82
+ )
83
+ if hasattr(os, "add_dll_directory"):
84
+ _DLL_HANDLES.append(os.add_dll_directory(str(r_bin)))
85
+ # R loads package DLLs itself. On current Windows/Python combinations,
86
+ # registering the directory with the OS is also required for their
87
+ # transitive dependencies (for example stats.dll -> R.dll).
88
+ if not ctypes.windll.kernel32.SetDllDirectoryW(str(r_bin)):
89
+ raise BackendUnavailableError(
90
+ f"Windows could not register R's DLL directory: {r_bin}"
91
+ )
92
+ # Explicitly retain R's core DLLs as well. R 4.6 can otherwise start
93
+ # successfully while later failing to load recommended packages such
94
+ # as stats because their Rblas/Rlapack dependencies are not found.
95
+ for dll_name in ("R.dll", "Rblas.dll", "Rlapack.dll", "Riconv.dll"):
96
+ dll_path = r_bin / dll_name
97
+ if dll_path.is_file():
98
+ try:
99
+ _DLL_HANDLES.append(ctypes.WinDLL(str(dll_path)))
100
+ except OSError as exc:
101
+ raise BackendUnavailableError(
102
+ f"Windows could not load R's core library: {dll_path}"
103
+ ) from exc
104
+
105
+ return r_home
106
+
107
+
108
+ class RBackend:
109
+ """Own the embedded-R session and keep fitted models opaque."""
110
+
111
+ def __init__(self) -> None:
112
+ self.r_home = _configure_r_process()
113
+ rpy2_situation = None
114
+ original_get_r_flags = None
115
+ try:
116
+ # rpy2 3.6.x probes ``R CMD config --ldflags`` before loading
117
+ # R.dll. Some Windows R builds return no stdout for that query;
118
+ # rpy2 then raises IndexError instead of taking its normal
119
+ # Windows DLL-directory fallback. Translate only that empty-output
120
+ # case into the exception the fallback already handles.
121
+ if os.name == "nt":
122
+ import rpy2.situation as rpy2_situation
123
+
124
+ original_get_r_flags = rpy2_situation.get_r_flags
125
+
126
+ def get_r_flags_with_windows_fallback(*args: Any, **kwargs: Any) -> Any:
127
+ try:
128
+ return original_get_r_flags(*args, **kwargs)
129
+ except IndexError as exc:
130
+ raise subprocess.CalledProcessError(1, "R CMD config") from exc
131
+
132
+ rpy2_situation.get_r_flags = get_r_flags_with_windows_fallback
133
+
134
+ import rpy2.robjects as ro
135
+ from rpy2.rinterface_lib import openrlib
136
+ from rpy2.robjects import default_converter, pandas2ri
137
+ from rpy2.robjects.packages import importr
138
+ except Exception as exc: # pragma: no cover - depends on host setup
139
+ raise BackendUnavailableError(
140
+ "rpy2 could not initialize embedded R. Run geds.diagnostics() "
141
+ "and verify that R and rpy2 use compatible versions."
142
+ ) from exc
143
+ finally:
144
+ if rpy2_situation is not None and original_get_r_flags is not None:
145
+ rpy2_situation.get_r_flags = original_get_r_flags
146
+
147
+ self.ro = ro
148
+ self._openrlib = openrlib
149
+ self._converter = default_converter + pandas2ri.converter
150
+
151
+ extra_library = os.environ.get("GEDS_R_LIBRARY")
152
+ if extra_library:
153
+ current = list(ro.r(".libPaths()"))
154
+ ro.r[".libPaths"](ro.StrVector([extra_library, *current]))
155
+
156
+ try:
157
+ self.geds = importr("GeDS")
158
+ self.stats = importr("stats")
159
+ except Exception as exc:
160
+ raise BackendUnavailableError(
161
+ "The GeDS R package could not be loaded. Install GeDS into a "
162
+ "library visible to this R installation."
163
+ ) from exc
164
+
165
+ self.geds_version = str(
166
+ self.ro.r("as.character(packageVersion('GeDS'))")[0]
167
+ )
168
+ is_supported = bool(
169
+ self.ro.r(
170
+ "packageVersion('GeDS') >= "
171
+ f"package_version('{MIN_GEDS_VERSION}')"
172
+ )[0]
173
+ )
174
+ if not is_supported:
175
+ raise BackendUnavailableError(
176
+ f"GeDS {self.geds_version} is installed, but geds-python "
177
+ f"requires GeDS >= {MIN_GEDS_VERSION}."
178
+ )
179
+ capabilities = self.ro.r(
180
+ "get0('.GeDS_python_bridge_capabilities', "
181
+ "envir=asNamespace('GeDS'), inherits=FALSE)"
182
+ )
183
+ available = set() if capabilities is self.ro.NULL else set(map(str, capabilities))
184
+ missing = REQUIRED_GEDS_CAPABILITIES - available
185
+ if missing:
186
+ raise BackendUnavailableError(
187
+ "This GeDS R build lacks Python bridge fixes "
188
+ f"({', '.join(sorted(missing))}). Install the GeDS GitHub "
189
+ "commit specified in the geds-python README."
190
+ )
191
+
192
+ @contextmanager
193
+ def locked(self) -> Iterator[None]:
194
+ """Serialize access to R's process-global runtime."""
195
+ with self._openrlib.rlock:
196
+ yield
197
+
198
+ def dataframe_to_r(self, frame: pd.DataFrame) -> Any:
199
+ with self.locked(), self._converter.context():
200
+ return self.ro.conversion.get_conversion().py2rpy(frame)
201
+
202
+ def vector(self, values: Any) -> Any:
203
+ return self.ro.FloatVector(np.asarray(values, dtype=float))
204
+
205
+ def formula(self, expression: str) -> Any:
206
+ return self.ro.Formula(expression)
207
+
208
+ def family(self, name: str, link: str | None) -> Any:
209
+ normalized = name.lower()
210
+ functions = {
211
+ "gaussian": "gaussian",
212
+ "poisson": "poisson",
213
+ "quasipoisson": "quasipoisson",
214
+ "binomial": "binomial",
215
+ "quasibinomial": "quasibinomial",
216
+ "gamma": "Gamma",
217
+ }
218
+ if normalized not in functions:
219
+ choices = ", ".join(sorted(functions))
220
+ raise ValueError(f"Unsupported family {name!r}; choose one of: {choices}.")
221
+ function = getattr(self.stats, functions[normalized])
222
+ return function() if link is None else function(link=link)
223
+
224
+ def boost_family(self, name: str) -> Any:
225
+ """Construct an mboost loss family for GeDS's R boosting algorithm."""
226
+ from rpy2.robjects.packages import importr
227
+
228
+ functions = {
229
+ "gaussian": "Gaussian",
230
+ "poisson": "Poisson",
231
+ "binomial": "Binomial",
232
+ "gamma": "GammaReg",
233
+ }
234
+ normalized = name.lower()
235
+ if normalized not in functions:
236
+ choices = ", ".join(sorted(functions))
237
+ raise ValueError(f"Unsupported boosting family {name!r}; choose one of: {choices}.")
238
+ with self.locked():
239
+ return getattr(importr("mboost"), functions[normalized])()
240
+
241
+ def fit(self, function: str, formula: Any, data: Any, **kwargs: Any) -> Any:
242
+ with self.locked():
243
+ return getattr(self.geds, function)(formula, data=data, **kwargs)
244
+
245
+ def predict(
246
+ self, model: Any, data: Any, order: int, prediction_type: str,
247
+ *, base_learner: str | None = None,
248
+ ) -> np.ndarray | pd.DataFrame:
249
+ with self.locked():
250
+ kwargs: dict[str, Any] = {"newdata": data, "n": order, "type": prediction_type}
251
+ if base_learner is not None:
252
+ kwargs["base_learner"] = base_learner
253
+ result = self.stats.predict(model, **kwargs)
254
+ values = np.array(result, dtype=float, copy=True)
255
+ if prediction_type == "terms":
256
+ names = self.ro.r("colnames")(result)
257
+ if values.ndim != 2 or names is self.ro.NULL:
258
+ raise BackendUnavailableError(
259
+ "This GeDS R build does not return named term predictions. "
260
+ "Install the GeDS GitHub prediction fixes."
261
+ )
262
+ columns = [str(name) for name in names]
263
+ return pd.DataFrame(values, columns=columns)
264
+ return values
265
+
266
+ def coefficients(self, model: Any, order: int) -> Any:
267
+ with self.locked():
268
+ return self.to_python(self.stats.coef(model, n=order))
269
+
270
+ def knots(self, model: Any, order: int) -> Any:
271
+ with self.locked():
272
+ return self.to_python(
273
+ self.stats.knots(model, n=order, options="internal")
274
+ )
275
+
276
+ def deviance(self, model: Any, order: int) -> float:
277
+ with self.locked():
278
+ return float(self.stats.deviance(model, n=order)[0])
279
+
280
+ def log_likelihood(self, model: Any, order: int) -> float:
281
+ with self.locked():
282
+ return float(self.stats.logLik(model, n=order)[0])
283
+
284
+ def confidence_intervals(
285
+ self, model: Any, order: int, level: float
286
+ ) -> pd.DataFrame:
287
+ with self.locked():
288
+ intervals = self.stats.confint(model, n=order, level=level)
289
+ names = [str(name) for name in self.ro.r("rownames")(intervals)]
290
+ values = np.array(intervals, dtype=float, copy=True)
291
+ return pd.DataFrame(values, index=names, columns=["lower", "upper"])
292
+
293
+ def derive(self, model: Any, x: Any, order: int, spline_order: int) -> np.ndarray:
294
+ with self.locked():
295
+ result = self.geds.Derive(
296
+ model, order=order, x=self.vector(x), n=spline_order
297
+ )
298
+ return np.array(result, dtype=float, copy=True)
299
+
300
+ def integrate(self, model: Any, lower: Any, upper: Any, spline_order: int) -> np.ndarray:
301
+ with self.locked():
302
+ result = self.geds.Integrate(
303
+ model, **{"from": self.vector(lower)}, to=self.vector(upper),
304
+ n=spline_order,
305
+ )
306
+ return np.array(result, dtype=float, copy=True)
307
+
308
+ def piecewise_polynomial(self, model: Any, spline_order: int) -> tuple[np.ndarray, np.ndarray]:
309
+ with self.locked():
310
+ result = self.geds.PPolyRep(model, n=spline_order)
311
+ knots = np.array(result.rx2("knots"), dtype=float, copy=True)
312
+ coefficients = np.array(result.rx2("coefficients"), dtype=float, copy=True)
313
+ return knots, coefficients
314
+
315
+ def shape_constrain(
316
+ self, model: Any, spline_order: int, constraints: list[str],
317
+ eps: float, ridge: float, base_learner: str | None,
318
+ ) -> Any:
319
+ with self.locked():
320
+ kwargs: dict[str, Any] = {
321
+ "n": spline_order,
322
+ "shape_constraint": self.ro.StrVector(constraints),
323
+ "eps": eps,
324
+ "ridge": ridge,
325
+ }
326
+ if base_learner is not None:
327
+ kwargs["base_learner"] = base_learner
328
+ return self.geds.shapeConstrain(model, **kwargs)
329
+
330
+ def base_learner_importance(
331
+ self, model: Any, boosting_iter_only: bool
332
+ ) -> pd.Series:
333
+ with self.locked():
334
+ result = self.geds.bl_imp(
335
+ model, boosting_iter_only=boosting_iter_only
336
+ )
337
+ names = [str(name) for name in self.ro.r("names")(result)]
338
+ values = np.array(result, dtype=float, copy=True)
339
+ return pd.Series(values, index=names, name="importance")
340
+
341
+ def cross_validate(
342
+ self, model_name: str, formula: str, frame: pd.DataFrame,
343
+ parameters: dict[str, np.ndarray], order: int, n_folds: int,
344
+ n_cores: int, random_state: int | None,
345
+ **fit_kwargs: Any,
346
+ ) -> tuple[pd.DataFrame, pd.DataFrame]:
347
+ """Call R's specialized GeDS grid search, retaining its model symbol."""
348
+ if model_name not in {"NGeDS", "GGeDS", "NGeDSgam", "NGeDSboost"}:
349
+ raise ValueError("Unsupported GeDS cross-validation model.")
350
+ with self.locked():
351
+ # crossv_GeDS inspects the *symbol* of model_fun via substitute().
352
+ # Passing an rpy2 function value directly would lose that symbol.
353
+ wrapper = self.ro.r(
354
+ "function(formula, data, parameters, n, n_folds, n_cores, ...) {"
355
+ f" {model_name} <- GeDS::{model_name};"
356
+ " GeDS::crossv_GeDS(formula=formula, data=data,"
357
+ f" model_fun={model_name}, parameters=parameters,"
358
+ " n=n, n_folds=n_folds, n_cores=n_cores, ...)"
359
+ " }"
360
+ )
361
+ r_parameters = self.ro.ListVector({
362
+ name: self.vector(values) for name, values in parameters.items()
363
+ })
364
+ if random_state is not None:
365
+ self.ro.r["set.seed"](int(random_state))
366
+ with self._converter.context():
367
+ r_frame = self.ro.conversion.get_conversion().py2rpy(frame)
368
+ library = os.environ.get("GEDS_R_LIBRARY")
369
+ original_r_libs_user = os.environ.get("R_LIBS_USER")
370
+ if library:
371
+ os.environ["R_LIBS_USER"] = os.pathsep.join(
372
+ part for part in (library, original_r_libs_user) if part
373
+ )
374
+ try:
375
+ result = wrapper(
376
+ self.formula(formula), r_frame,
377
+ parameters=r_parameters, n=order, n_folds=n_folds,
378
+ n_cores=n_cores, **fit_kwargs,
379
+ )
380
+ finally:
381
+ if library:
382
+ if original_r_libs_user is None:
383
+ os.environ.pop("R_LIBS_USER", None)
384
+ else:
385
+ os.environ["R_LIBS_USER"] = original_r_libs_user
386
+ with self._converter.context():
387
+ conversion = self.ro.conversion.get_conversion()
388
+ best = conversion.rpy2py(result.rx2("best_params"))
389
+ results = conversion.rpy2py(result.rx2("results"))
390
+ return pd.DataFrame(best).reset_index(drop=True), pd.DataFrame(results).reset_index(drop=True)
391
+
392
+ def save_boosting_diagnostics(
393
+ self, model: Any, path: Path, iterations: list[int], final_fits: bool
394
+ ) -> None:
395
+ """Write R ``visualize_boosting`` plots to a multipage PDF."""
396
+ with self.locked():
397
+ self.ro.r["pdf"](file=str(path), width=8, height=6, onefile=True)
398
+ try:
399
+ self.geds.visualize_boosting(
400
+ model, iters=self.ro.IntVector(iterations),
401
+ final_fits=final_fits,
402
+ )
403
+ finally:
404
+ self.ro.r["dev.off"]()
405
+
406
+ def component(self, model: Any, name: str) -> Any:
407
+ try:
408
+ return self.to_python(model.rx2(name))
409
+ except Exception:
410
+ return None
411
+
412
+ def to_python(self, value: Any) -> Any:
413
+ """Convert extracted values without recursively converting a model."""
414
+ vectors = self.ro.vectors
415
+ if value is self.ro.NULL:
416
+ return None
417
+ if isinstance(value, vectors.ListVector):
418
+ names = list(value.names) if value.names is not self.ro.NULL else []
419
+ converted = [self.to_python(item) for item in value]
420
+ if names and len(names) == len(converted) and all(names):
421
+ return dict(zip(names, converted))
422
+ return converted
423
+ if isinstance(value, vectors.StrVector):
424
+ result = np.array(value, dtype=str, copy=True)
425
+ elif isinstance(
426
+ value, (vectors.FloatVector, vectors.IntVector, vectors.BoolVector)
427
+ ):
428
+ result = np.array(value, copy=True)
429
+ else:
430
+ return value
431
+ return result.item() if result.ndim == 0 else result
432
+
433
+ def information(self) -> dict[str, str]:
434
+ with self.locked():
435
+ return {
436
+ "r_home": str(self.r_home),
437
+ "r_version": str(self.ro.r("R.version.string")[0]),
438
+ "geds_version": self.geds_version,
439
+ "minimum_geds_version": MIN_GEDS_VERSION,
440
+ "geds_library": str(self.ro.r("find.package('GeDS')")[0]),
441
+ "rpy2_version": version("rpy2"),
442
+ "python": sys.version.split()[0],
443
+ }
444
+
445
+
446
+ _BACKEND: RBackend | None = None
447
+
448
+
449
+ def get_backend() -> RBackend:
450
+ global _BACKEND
451
+ if _BACKEND is None:
452
+ _BACKEND = RBackend()
453
+ return _BACKEND
454
+
455
+
456
+ def diagnostics() -> dict[str, str]:
457
+ """Return the versions and locations used by the Python/R bridge."""
458
+ return get_backend().information()