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 +25 -0
- geds/_backend.py +458 -0
- geds/_estimators.py +844 -0
- geds/_plotting.py +69 -0
- geds/_validation.py +148 -0
- geds/check.py +64 -0
- geds/py.typed +1 -0
- geds_python-0.1.0.dist-info/METADATA +406 -0
- geds_python-0.1.0.dist-info/RECORD +11 -0
- geds_python-0.1.0.dist-info/WHEEL +4 -0
- geds_python-0.1.0.dist-info/licenses/LICENSE +674 -0
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()
|