Skip to content

voiage.backends.advanced_jax_regression.JaxAdvancedRegression

Advanced JAX-optimized regression models for EVPPI calculations.

polynomial_features([positional or keyword] self: None = None, [positional or keyword] x: np.ndarray = None, [positional or keyword] degree: int = 2) -> np.ndarray

Generate polynomial features for regression.

Parameters:

  • self
  • x np.ndarray
  • degree int (default: 2)

Returns: np.ndarray

fit_polynomial([positional or keyword] self: None = None, [positional or keyword] x: np.ndarray = None, [positional or keyword] y: np.ndarray = None, [positional or keyword] degree: int = 2, [positional or keyword] regularization: float = 1e-06) -> JaxAdvancedRegression

Fit polynomial regression using JAX optimization.

Parameters:

  • self
  • x np.ndarray
  • y np.ndarray
  • degree int (default: 2)
  • regularization float (default: 1e-06)

Returns: JaxAdvancedRegression

predict([positional or keyword] self: None = None, [positional or keyword] x: np.ndarray = None) -> np.ndarray

Make predictions using fitted model.

Parameters:

  • self
  • x np.ndarray

Returns: np.ndarray

r_squared([positional or keyword] self: None = None, [positional or keyword] x: np.ndarray = None, [positional or keyword] y: np.ndarray = None) -> float

Calculate R-squared score.

Parameters:

  • self
  • x np.ndarray
  • y np.ndarray

Returns: float

cross_validate([positional or keyword] self: None = None, [positional or keyword] x: np.ndarray = None, [positional or keyword] y: np.ndarray = None, [positional or keyword] degree: int = 2, [positional or keyword] n_folds: int = 5, [positional or keyword] regularization: float = 1e-06) -> float

Perform cross-validation to find optimal degree.

Parameters:

  • self
  • x np.ndarray
  • y np.ndarray
  • degree int (default: 2)
  • n_folds int (default: 5)
  • regularization float (default: 1e-06)

Returns: float