diff --git a/README.md b/README.md index 37f2d193..bc474098 100644 --- a/README.md +++ b/README.md @@ -125,7 +125,7 @@ result, and prediction interfaces. | Formula surface | Supported terms | | ------------------ | ---------------------------------------------------------------------------------------------------- | -| Univariate smooths | `s(..., bs='cr')`, `cs`, `cc`, `ps`, `tp`, `ts` | +| Univariate smooths | `s(..., bs='cr')`, `cs`, `cc`, `cp`, `ps`, `tp`, `ts` | | Structured smooths | random effects `re`, factor smooths `fs`, sum-to-zero factor smooths `sz` | | Tensor products | `te(...)` and `ti(...)` over supported numeric marginals | | Parametric terms | numeric and factor terms, supported interactions, intercept policies, and formula offsets | diff --git a/nampy/gam/compiler/factory.py b/nampy/gam/compiler/factory.py index c3d1fece..f82708fa 100644 --- a/nampy/gam/compiler/factory.py +++ b/nampy/gam/compiler/factory.py @@ -139,14 +139,15 @@ def instantiate_term(term_like: TermSpec | Any): ) if isinstance(smooth_spec, PSplineSmoothSpec): + basis = str(smooth_spec.bs).lower() if len(features) != 1: raise NotImplementedError( - "Current runtime only materializes 1D s(..., bs='ps') terms." + f"Current runtime only materializes 1D s(..., bs={basis!r}) terms." ) return PSplineTerm1D( feature=features[0], k=smooth_spec.k, - basis="ps", + basis=basis, m=smooth_spec.m, label=label, term_id=term_like.term_id, diff --git a/nampy/gam/smooths/__init__.py b/nampy/gam/smooths/__init__.py index 97559401..1fc77f41 100644 --- a/nampy/gam/smooths/__init__.py +++ b/nampy/gam/smooths/__init__.py @@ -29,7 +29,7 @@ ti = InteractionTensorProductSplineTerm cr = cs = cc = CubicSplineTerm -ps = PSplineTerm1D +cp = ps = PSplineTerm1D tp = ts = ThinPlateSplineTerm fs = FSmoothInteractionTerm sz = SZSmoothInteractionTerm @@ -68,6 +68,7 @@ "cr", "cs", "cc", + "cp", "ps", "tp", "ts", diff --git a/nampy/gam/smooths/categorical/fs.py b/nampy/gam/smooths/categorical/fs.py index 8c78100c..414e455f 100644 --- a/nampy/gam/smooths/categorical/fs.py +++ b/nampy/gam/smooths/categorical/fs.py @@ -106,7 +106,7 @@ def _build_base_smooth_term( Build the per-level base smooth used inside fs/sz. Supported base smooth classes in the current codebase: - cr, cs, cc, ps, tp, ts + cr, cs, cc, cp, ps, tp, ts """ base_bs = str(base_bs).lower() metric_features = list(metric_features) @@ -123,9 +123,10 @@ def _build_base_smooth_term( f"for bs in {{'tp','ts'}}, got base bs={base_bs!r}." ) - if xt_rest is not None and base_bs not in {"tp", "ts", "ps"}: + if xt_rest is not None and base_bs not in {"tp", "ts", "ps", "cp"}: raise NotImplementedError( - f"Extra xt options are currently only supported for tp/ts/ps base smooths, " + "Extra xt options are currently only supported for tp/ts/ps/cp base " + "smooths, " f"got xt={xt_rest!r} with base bs={base_bs!r}." ) @@ -147,7 +148,7 @@ def _build_base_smooth_term( metadata=metadata, ) - if base_bs == "ps": + if base_bs in {"ps", "cp"}: ps_m = None if xt_rest is None else xt_rest.get("m", None) # For fs/sz, mgcv keeps the outer basis dimension and uses xt mainly to # choose the base smoother family / order parameters. @@ -155,7 +156,7 @@ def _build_base_smooth_term( return PSplineTerm1D( feature=metric_features[0], k=ps_k, - basis="ps", + basis=base_bs, m=ps_m, label=label, smoothing_id=None, @@ -191,12 +192,14 @@ def _build_base_smooth_term( raise NotImplementedError( f"Current {mode} implementation supports base bs in " - f"{{'cr','cs','cc','ps','tp','ts'}}, got {base_bs!r}." + f"{{'cr','cs','cc','cp','ps','tp','ts'}}, got {base_bs!r}." ) def _penalty_rank_from_base_term(base_term, basis_matrix, penalty_matrix) -> int: if isinstance(base_term, PSplineTerm1D) and len(base_term.penalties) > 0: + if str(base_term.basis_name).lower() == "cp": + return int(base_term._setup.rank) # mgcv::smooth.construct.ps.smooth.spec uses rank <- bs.dim - m[2]. penalty_order = int(base_term.m[1]) return max(0, int(basis_matrix.shape[1]) - penalty_order) diff --git a/nampy/gam/smooths/tensor/marginals.py b/nampy/gam/smooths/tensor/marginals.py index 281843c3..dc11c7ce 100644 --- a/nampy/gam/smooths/tensor/marginals.py +++ b/nampy/gam/smooths/tensor/marginals.py @@ -11,7 +11,7 @@ from ..univariate.ps import PSplineTerm1D from ..univariate.tp import ThinPlateSplineTerm -TENSOR_MARGINAL_BASES = frozenset({"cr", "cs", "cc", "ps", "tp", "ts"}) +TENSOR_MARGINAL_BASES = frozenset({"cr", "cs", "cc", "cp", "ps", "tp", "ts"}) def _as_marginal_features(feature): @@ -72,11 +72,11 @@ def make_tensor_marginal_term( metadata=metadata, ) - if basis == "ps": + if basis in {"ps", "cp"}: if len(marginal_features) != 1: raise ValueError( - "Tensor marginal basis 'ps' only handles one feature; mgcv coerces " - "multivariate ps marginals to tp before construction." + f"Tensor marginal basis {basis!r} only handles one feature; mgcv " + "coerces multivariate ps/cp marginals to tp before construction." ) return PSplineTerm1D( feature=marginal_features[0], diff --git a/nampy/gam/smooths/univariate/__init__.py b/nampy/gam/smooths/univariate/__init__.py index d30969f0..3021ca13 100644 --- a/nampy/gam/smooths/univariate/__init__.py +++ b/nampy/gam/smooths/univariate/__init__.py @@ -3,8 +3,8 @@ from .tp import ThinPlateSplineTerm cr = cs = cc = CubicSplineTerm -ps = PSplineTerm1D +cp = ps = PSplineTerm1D tp = ts = ThinPlateSplineTerm __all__ = ["CubicSplineTerm", "PSplineTerm1D", "ThinPlateSplineTerm"] -__all__ += ["cr", "cs", "cc", "ps", "tp", "ts"] +__all__ += ["cr", "cs", "cc", "cp", "ps", "tp", "ts"] diff --git a/nampy/gam/smooths/univariate/ps.py b/nampy/gam/smooths/univariate/ps.py index 48796c48..4d1f534c 100644 --- a/nampy/gam/smooths/univariate/ps.py +++ b/nampy/gam/smooths/univariate/ps.py @@ -1,11 +1,9 @@ """ -P-spline smooth term (``bs='ps'``). +P-spline smooth terms (``bs='ps'`` and cyclic ``bs='cp'``). Implements the :class:`BaseSmoothTerm` interface for a P-spline: a B-spline -basis with a discrete difference penalty on adjacent coefficients. Unlike -regression splines, P-splines do not require a set of knots to be chosen -ahead of time; instead, many equally-spaced knots are used and the smoothness -is controlled entirely by the penalty order and the smoothing parameter. +basis with a discrete coefficient-difference penalty. Cyclic P-splines use a +wrapped basis, a circular difference penalty, and periodic newdata mapping. """ import numpy as np @@ -15,7 +13,7 @@ from ...splines.univariate.ps import ( build_pspline_term_setup, predict_pspline_term, - pspline_predict_matrix, + predict_pspline_term_derivative, ) from ..registry import register_smooth from ..smooth_base import ( @@ -29,6 +27,7 @@ @register_smooth("ps") +@register_smooth("cp") class PSplineTerm1D(BaseSmoothTerm): term_type = "smooth" basis_name = "ps" @@ -71,22 +70,36 @@ def __init__( self.knots = knots self.null_penalty_tol = float(null_penalty_tol) + def normalize_order(value): + if value is None or (isinstance(value, float) and np.isnan(value)): + return 2 + numeric = float(value) + if not np.isfinite(numeric) or numeric != np.rint(numeric): + raise ValueError( + f"For bs={self.basis_name!r}, m entries must be integers or NA." + ) + return int(numeric) + if m is None: self.m = (2, 2) elif np.isscalar(m): - self.m = (int(m), int(m)) + value = normalize_order(m) + self.m = (value, value) else: - vals = tuple(int(v) for v in m) + vals = tuple(normalize_order(v) for v in m) if len(vals) == 1: self.m = (vals[0], vals[0]) - elif len(vals) == 2: + elif len(vals) == 2 or (self.basis_name == "cp" and len(vals) > 2): self.m = vals else: - raise ValueError("For bs='ps', m must have length 1 or 2.") + raise ValueError( + f"For bs={self.basis_name!r}, m must have length 1 or 2." + ) - if self.basis_name != "ps": + if self.basis_name not in {"ps", "cp"}: raise NotImplementedError( - f"PSplineTerm1D currently supports only basis='ps', got {basis!r}." + "PSplineTerm1D supports only basis in {'ps', 'cp'}, " + f"got {basis!r}." ) if self.select and self.fixed: raise ValueError("select=True and fixed=True are incompatible.") @@ -131,9 +144,11 @@ def fit(self, X, feature_names): else: x_setup_values = np.asarray(xj, dtype=np.float64).reshape(-1) - basis_order, penalty_order = self.m + basis_order, penalty_order = self.m[:2] if basis_order < 0 or penalty_order < 0: - raise ValueError("For bs='ps', m entries must be >= 0.") + raise ValueError( + f"For bs={self.basis_name!r}, m entries must be >= 0." + ) shared_X = self._linked_id_setup_matrix(feature_names) if shared_X is not None: @@ -149,6 +164,7 @@ def fit(self, X, feature_names): bs_dim=self.k, m=self.m, knots=self.knots, + basis=self.basis_name, ) setup_base = np.asarray(self._setup.basis_train, dtype=np.float64) base = np.asarray(predict_pspline_term(xj, self._setup), dtype=np.float64) @@ -163,6 +179,7 @@ def fit(self, X, feature_names): bs_dim=self.k, m=self.m, knots=self.knots, + basis=self.basis_name, ) point_base = np.asarray(self._setup.basis_train, dtype=np.float64) if self._linear_functional: @@ -309,12 +326,7 @@ def derivative_matrix(self, X_new=None, order=1): ) source = self._X_train if X_new is None else X_new xj = column_as_numeric_array(source, self._feature_index) - B = pspline_predict_matrix( - xj, - self._setup.knots, - basis_order=self._setup.basis_order, - deriv=order, - ) + B = predict_pspline_term_derivative(xj, self._setup, deriv=order) return self._apply_constraint_transform_and_by(B, source) def tensor_marginal_fit_matrices( diff --git a/nampy/gam/specs/modeling.py b/nampy/gam/specs/modeling.py index bdffc483..feee207f 100644 --- a/nampy/gam/specs/modeling.py +++ b/nampy/gam/specs/modeling.py @@ -42,7 +42,7 @@ def make_predictor_specs(model, feature_names, *, knots=None): metadata={}, ) ) - elif basis == "ps": + elif basis in {"ps", "cp"}: main_terms.append( TermSpec( kind="smooth", @@ -50,7 +50,7 @@ def make_predictor_specs(model, feature_names, *, knots=None): by_variable=None, smooth_spec=build_smooth_spec( special="s", - bs="ps", + bs=basis, k=model.k, m=None, sp=None, @@ -105,7 +105,7 @@ def make_predictor_specs(model, feature_names, *, knots=None): else: raise NotImplementedError( "Automatic main-effect construction currently supports " - "{'cr','cs','cc','ps','tp','ts','re'}, " + "{'cr','cs','cc','cp','ps','tp','ts','re'}, " f"got {model.basis!r}." ) diff --git a/nampy/gam/specs/smooth_build.py b/nampy/gam/specs/smooth_build.py index c40aa68a..d75339de 100644 --- a/nampy/gam/specs/smooth_build.py +++ b/nampy/gam/specs/smooth_build.py @@ -86,6 +86,7 @@ def _build_s_cc(opts) -> CyclicCubicRegressionSmoothSpec: def _build_s_ps(opts) -> PSplineSmoothSpec: return PSplineSmoothSpec( special="s", + bs=str(opts["bs"]).lower(), k=opts["k"], fx=opts["fx"], select=opts["select"], @@ -182,6 +183,7 @@ def _build_s_sz(opts) -> SumToZeroFactorSmoothSpec: "cs": _build_s_cs, "cc": _build_s_cc, "ps": _build_s_ps, + "cp": _build_s_ps, "tp": _build_s_tp, "ts": _build_s_ts, "re": _build_s_re, @@ -274,7 +276,7 @@ def _is_vector_fx(fx) -> bool: return fx is not None and not np.isscalar(fx) -_PC_SUPPORTED_S_BASES = {"cc", "cr", "cs", "ps", "tp", "ts"} +_PC_SUPPORTED_S_BASES = {"cc", "cp", "cr", "cs", "ps", "tp", "ts"} def _dispatch_smooth_spec_from_options(opts) -> SmoothSpec: @@ -290,7 +292,7 @@ def _dispatch_smooth_spec_from_options(opts) -> SmoothSpec: raise NotImplementedError( f"pc= is not supported for s(..., bs={merged['bs']!r}); " "point constraints are only supported for bs in " - "{'cc', 'cr', 'cs', 'ps', 'tp', 'ts'}." + "{'cc', 'cp', 'cr', 'cs', 'ps', 'tp', 'ts'}." ) return builder(merged) if has_pc and special_key not in {"te", "ti"}: diff --git a/nampy/gam/splines/univariate/__init__.py b/nampy/gam/splines/univariate/__init__.py index ab468d80..65e53898 100644 --- a/nampy/gam/splines/univariate/__init__.py +++ b/nampy/gam/splines/univariate/__init__.py @@ -11,7 +11,12 @@ PSplineBasisSetup, bspline_design_matrix, build_pspline_term_setup, + cyclic_pspline_design, + cyclic_pspline_difference_penalty, + cyclic_pspline_knots, + cyclic_wrap, predict_pspline_term, + predict_pspline_term_derivative, pspline_difference_penalty, pspline_knots, pspline_predict_matrix, @@ -23,6 +28,10 @@ "bspline_design_matrix", "cyclic_cubic_bd", "cyclic_cubic_predict_matrix", + "cyclic_pspline_design", + "cyclic_pspline_difference_penalty", + "cyclic_pspline_knots", + "cyclic_wrap", "place_knots_through_values", "pspline_difference_penalty", "pspline_knots", @@ -31,6 +40,7 @@ "PSplineBasisSetup", "build_pspline_term_setup", "predict_pspline_term", + "predict_pspline_term_derivative", "build_tprs_term_setup", "predict_tprs_term", ] diff --git a/nampy/gam/splines/univariate/ps.py b/nampy/gam/splines/univariate/ps.py index 26e926e5..6e38580a 100644 --- a/nampy/gam/splines/univariate/ps.py +++ b/nampy/gam/splines/univariate/ps.py @@ -1,3 +1,4 @@ +import warnings from dataclasses import dataclass import numpy as np @@ -108,6 +109,128 @@ def pspline_difference_penalty(n_coef, diff_order): return D.T @ D +def cyclic_pspline_knots(x, bs_dim, basis_order, supplied_knots=None): + """Port ``smooth.construct.cp.smooth.spec`` knot handling.""" + x = np.asarray(x, dtype=np.float64).ravel() + nk = int(bs_dim) + 1 + if nk <= int(basis_order): + raise ValueError("basis dimension too small for b-spline order") + + if supplied_knots is None: + lower = float(np.min(x)) + upper = float(np.max(x)) + return np.linspace(lower, upper, nk) + + knots = np.asarray(supplied_knots, dtype=np.float64).ravel() + if knots.size == 2: + lower = float(np.min(knots)) + upper = float(np.max(knots)) + if lower > np.min(x) or upper < np.max(x): + raise ValueError("knot range does not include data") + return np.linspace(lower, upper, nk) + if knots.size != nk: + raise ValueError(f"there should be {nk} supplied knots") + return knots + + +def cyclic_wrap(x, lower, upper): + """Port mgcv's ``cwrap`` mapping onto a cyclic interval.""" + values = np.asarray(x, dtype=np.float64).ravel().copy() + lower = float(lower) + upper = float(upper) + width = upper - lower + above = values > upper + if np.any(above): + values[above] = lower + np.mod(values[above] - upper, width) + below = values < lower + if np.any(below): + values[below] = upper - np.mod(lower - values[below], width) + return values + + +def _outer_ok_bspline_design(x, knots, degree, deriv=0): + """Match ``splines::splineDesign(..., outer.ok=TRUE)`` on all knot spans.""" + values = np.asarray(x, dtype=np.float64).ravel() + knot_vector = np.asarray(knots, dtype=np.float64).ravel() + degree = int(degree) + deriv = int(deriv) + n_basis = int(knot_vector.size - degree - 1) + if deriv > degree: + return np.zeros((values.size, n_basis), dtype=np.float64) + design = np.zeros((values.size, n_basis), dtype=np.float64) + for index in range(n_basis): + basis = BSpline.basis_element( + knot_vector[index : index + degree + 2], extrapolate=False + ) + if deriv: + basis = basis.derivative(deriv) + design[:, index] = np.nan_to_num( + basis(values), nan=0.0, posinf=0.0, neginf=0.0 + ) + return design + + +def cyclic_pspline_design(x, knots, basis_order, deriv=0, *, wrap=False): + """Operation-for-operation port of mgcv's ``cSplineDes``.""" + values = np.asarray(x, dtype=np.float64).ravel().copy() + cyclic_knots = np.sort(np.asarray(knots, dtype=np.float64).ravel()) + order = int(basis_order) + 2 + deriv = int(deriv) + if order < 2: + raise ValueError("order too low") + if cyclic_knots.size < order: + raise ValueError("too few knots") + + lower = float(cyclic_knots[0]) + upper = float(cyclic_knots[-1]) + if wrap and (np.min(values) < lower or np.max(values) > upper): + values = cyclic_wrap(values, lower, upper) + if np.min(values) < lower or np.max(values) > upper: + raise ValueError("x out of range") + + wrap_threshold = float(cyclic_knots[cyclic_knots.size - order]) + prefix = lower - ( + upper - cyclic_knots[cyclic_knots.size - order : -1] + ) + extended_knots = np.concatenate([prefix, cyclic_knots]) + design = _outer_ok_bspline_design( + values, + extended_knots, + degree=order - 1, + deriv=deriv, + ) + wrapped_rows = values > wrap_threshold + if np.any(wrapped_rows): + shifted = values[wrapped_rows] - upper + lower + design[wrapped_rows, :] += _outer_ok_bspline_design( + shifted, + extended_knots, + degree=order - 1, + deriv=deriv, + ) + return np.asarray(design, dtype=np.float64) + + +def cyclic_pspline_difference_penalty(n_coef, diff_order): + """Port the wrapped difference penalty in the cyclic P-spline constructor.""" + n_coef = int(n_coef) + diff_order = int(diff_order) + if diff_order < 0: + raise ValueError("diff_order must be >= 0") + if diff_order > n_coef - 1: + raise ValueError("penalty order too high for basis dimension") + + expanded = np.eye(n_coef + diff_order, dtype=np.float64) + for _ in range(diff_order): + expanded = np.diff(expanded, axis=0) + if diff_order == 0: + difference = expanded + else: + difference = expanded[:, diff_order:].copy() + difference[:, n_coef - diff_order : n_coef] += expanded[:, :diff_order] + return difference.T @ difference + + def pspline_predict_matrix(x, knots, basis_order, deriv=0): """ mgcv::Predict.matrix.pspline.smooth analogue. @@ -203,6 +326,8 @@ class PSplineBasisSetup: penalty: np.ndarray bs_dim: int rank: int + basis_name: str = "ps" + orders: tuple[int, ...] = () def build_pspline_term_setup( @@ -213,27 +338,49 @@ def build_pspline_term_setup( bs_dim, m, knots=None, + basis="ps", ): x = np.asarray(x, dtype=np.float64).ravel() basis_order, penalty_order = (int(m[0]), int(m[1])) + basis_name = str(basis).lower() + if basis_name not in {"ps", "cp"}: + raise ValueError(f"Unsupported P-spline basis {basis!r}.") if basis_order < 0 or penalty_order < 0: - raise ValueError("For bs='ps', m entries must be >= 0.") - - k = pspline_knots( - x, - bs_dim=int(bs_dim), - basis_order=basis_order, - supplied_knots=knots, - ) - degree = basis_order + 1 - B = bspline_design_matrix( - x, - k, - degree=degree, - deriv=0, - extrapolate=True, - ) - S = pspline_difference_penalty(B.shape[1], penalty_order) + raise ValueError(f"For bs={basis_name!r}, m entries must be >= 0.") + + if basis_name == "cp": + k = cyclic_pspline_knots( + x, + bs_dim=int(bs_dim), + basis_order=basis_order, + supplied_knots=knots, + ) + B = cyclic_pspline_design(x, k, basis_order, deriv=0) + if np.any(np.sum(B, axis=0) == 0.0): + warnings.warn( + "knot range is so wide that there is *no* information about some " + "basis coefficients", + stacklevel=2, + ) + S = cyclic_pspline_difference_penalty(B.shape[1], penalty_order) + rank = int(B.shape[1] - 1) + else: + k = pspline_knots( + x, + bs_dim=int(bs_dim), + basis_order=basis_order, + supplied_knots=knots, + ) + degree = basis_order + 1 + B = bspline_design_matrix( + x, + k, + degree=degree, + deriv=0, + extrapolate=True, + ) + S = pspline_difference_penalty(B.shape[1], penalty_order) + rank = numerical_rank(S, hermitian=True) S = symmetrize_matrix(S) return PSplineBasisSetup( @@ -245,12 +392,22 @@ def build_pspline_term_setup( basis_train=np.asarray(B, dtype=np.float64), penalty=np.asarray(S, dtype=np.float64), bs_dim=int(B.shape[1]), - rank=numerical_rank(S, hermitian=True), + rank=rank, + basis_name=basis_name, + orders=tuple(int(value) for value in m), ) def predict_pspline_term(x_new, setup: PSplineBasisSetup): x_new = np.asarray(x_new, dtype=np.float64).ravel() + if str(setup.basis_name).lower() == "cp": + return cyclic_pspline_design( + x_new, + setup.knots, + basis_order=setup.basis_order, + deriv=0, + wrap=True, + ) return np.asarray( pspline_predict_matrix( x_new, @@ -260,3 +417,22 @@ def predict_pspline_term(x_new, setup: PSplineBasisSetup): ), dtype=np.float64, ) + + +def predict_pspline_term_derivative(x_new, setup: PSplineBasisSetup, deriv=1): + """Evaluate an ordinary or cyclic P-spline derivative matrix.""" + x_new = np.asarray(x_new, dtype=np.float64).ravel() + if str(setup.basis_name).lower() == "cp": + return cyclic_pspline_design( + x_new, + setup.knots, + basis_order=setup.basis_order, + deriv=deriv, + wrap=True, + ) + return pspline_predict_matrix( + x_new, + setup.knots, + basis_order=setup.basis_order, + deriv=deriv, + ) diff --git a/pyproject.toml b/pyproject.toml index 322d9a92..9b18482a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -205,6 +205,7 @@ markers = [ "smooth_cr: tests covering cubic regression smooths", "smooth_cs: tests covering shrinkage cubic regression smooths", "smooth_cc: tests covering cyclic cubic smooths", + "smooth_cp: tests covering cyclic P-spline smooths", "smooth_ps: tests covering P-spline smooths", "smooth_tp: tests covering thin plate smooths", "smooth_ts: tests covering shrinkage thin plate smooths", diff --git a/tests/SUBSYSTEM_COVERAGE.md b/tests/SUBSYSTEM_COVERAGE.md index 717416e9..86b2501e 100644 --- a/tests/SUBSYSTEM_COVERAGE.md +++ b/tests/SUBSYSTEM_COVERAGE.md @@ -15,7 +15,7 @@ local development notes rather than duplicated here. | Subsystem | Primary owner(s) | Primary tests | Notes | | --- | --- | --- | --- | | Formula/spec parsing | `nampy/gam/formula/`, `nampy/gam/specs/` | `tests/parity/test_mgcv_formula_parse_parity.py` | Direct formula parity vs `mgcv`. | -| Smooth constructors / raw basis owners | `nampy/gam/smooths/`, `nampy/gam/splines/` | `tests/smooths/test_mgcv_raw_constructor_parity.py`, `tests/smooths/test_mgcv_smoothcon_parity.py`, `tests/parity/test_gam_spec_build_owner_contracts.py` | Basis, penalty, constructor, and smoothCon surfaces, including the upstream tensor-`m` wrong-length warning and zero fallback. | +| Smooth constructors / raw basis owners | `nampy/gam/smooths/`, `nampy/gam/splines/` | `tests/smooths/test_mgcv_raw_constructor_parity.py`, `tests/smooths/test_mgcv_smoothcon_parity.py`, `tests/parity/test_gam_spec_build_owner_contracts.py`, `tests/parity/test_mgcv_cp_combinations_parity.py` | Basis, penalty, constructor, and smoothCon surfaces, including cyclic P-splines (`cp`) across wrapped prediction, tensor/factor-smooth combinations, and the upstream tensor-`m` wrong-length warning and zero fallback. | | `pc=` / linked `id=` routing | smooth metadata + linked basis owners | `tests/smooths/test_mgcv_pc_id_parity.py` | Localizes shared-smoothing and point-constraint issues, including `te`/`ti` point constraints. | | Design / pre-fit assembly | `nampy/gam/compiler/`, `nampy/gam/fit/penalized_system.py` | `tests/optimization/test_mgcv_gam_setup_assembly_parity.py`, `tests/optimization/test_mgcv_preoptimization_blocks_parity.py`, `tests/optimization/test_mgcv_preoptimization_reparam_parity.py` | Setup, blocks, and reparameterization parity, including one global shared-component block with overlapping linear-predictor indices. | | Term wrapping / by-variable / offset routing | predictor wrapping + compiled term owners | `tests/optimization/test_gam_term_wrapping_owner_contracts.py` | Localizes wrapped predictor blocks, offset routing, and general-family block ownership before broader prediction parity. | diff --git a/tests/TAXONOMY.md b/tests/TAXONOMY.md index 6e969539..8694bedc 100644 --- a/tests/TAXONOMY.md +++ b/tests/TAXONOMY.md @@ -12,7 +12,7 @@ The GAM test suite is intentionally overlapping. The goal is fast subset runs an - `tests/`: shared helpers, marker inference, taxonomy registry, static reference fixtures, and parity-generation R scripts ## Taxonomy Axes -- `smooth_`: `cr`, `cs`, `cc`, `ps`, `tp`, `ts`, `te`, `ti`, `fs`, `sz`, `re` +- `smooth_`: `cr`, `cs`, `cc`, `cp`, `ps`, `tp`, `ts`, `te`, `ti`, `fs`, `sz`, `re` - `family_`: `gaussian`, `binomial`, `poisson`, `gamma`, `negbin`, `gaulss`, `gammals`, `general` - `method_`: `fixed`, `reml`, `ml`, `laml`, `gcv`, `ubre` - `link_`: `identity`, `log`, `inverse`, `logit`, `probit`, `cloglog`, `cauchit`, `sqrt` diff --git a/tests/_taxonomy_registry.py b/tests/_taxonomy_registry.py index 01368c27..0ca0c388 100644 --- a/tests/_taxonomy_registry.py +++ b/tests/_taxonomy_registry.py @@ -17,6 +17,7 @@ def _leaf(leaf_id: str, *nodeid_parts: str) -> LeafCoverageExpectation: "cr": "smooth_cr", "cs": "smooth_cs", "cc": "smooth_cc", + "cp": "smooth_cp", "ps": "smooth_ps", "tp": "smooth_tp", "ts": "smooth_ts", @@ -136,6 +137,11 @@ def _leaf(leaf_id: str, *nodeid_parts: str) -> LeafCoverageExpectation: "tests/smooths/test_mgcv_pc_id_parity.py", "tests/smooths/test_mgcv_smoothcon_parity.py", ), + "smooth_cp": ( + "tests/parity/test_mgcv_cp_combinations_parity.py", + "tests/smooths/test_mgcv_raw_constructor_parity.py", + "tests/smooths/test_mgcv_smoothcon_parity.py", + ), "smooth_ps": ( "tests/smooths/test_mgcv_raw_constructor_parity.py", "tests/smooths/test_mgcv_smoothcon_parity.py", diff --git a/tests/diagnostics/test_mgcv_k_check_parity.py b/tests/diagnostics/test_mgcv_k_check_parity.py index 44d83852..15d8d321 100644 --- a/tests/diagnostics/test_mgcv_k_check_parity.py +++ b/tests/diagnostics/test_mgcv_k_check_parity.py @@ -261,6 +261,14 @@ class TestKCheckParity: {"x0", "x1"}, 1e-4, ), + ( + lambda: _make_gaussian_data(seed=601, n=180), + 'y ~ s(x0, bs="cp", k=8) + s(x1, bs="cp", k=8)', + "gaussian", + "REML", + {"x0", "x1"}, + 1e-4, + ), ( lambda: _make_gaussian_data(seed=123, n=180), 'y ~ te(x0, x1, bs=["cr","cr"], k=[5,5])', @@ -274,6 +282,7 @@ class TestKCheckParity: "gaussian_cr", "gaussian_cr_fixed", "gaussian_ps", + "gaussian_cp", "gaussian_te", ], ) diff --git a/tests/mgcv_parity_utils.py b/tests/mgcv_parity_utils.py index 0d92fc99..48eb9ce8 100644 --- a/tests/mgcv_parity_utils.py +++ b/tests/mgcv_parity_utils.py @@ -1604,6 +1604,10 @@ def _run_mgcv_raw_constructor( knots = pack_vector(sm$knots, "numeric"), m = pack_vector(sm$m, "integer") ), + "cpspline.smooth" = list( + knots = pack_vector(sm$knots, "numeric"), + m = pack_vector(sm$m, "integer") + ), "tprs.smooth" = list( Xu = pack_matrix(sm$Xu), UZ = pack_matrix(sm$UZ), diff --git a/tests/parity/test_gam_pipeline_combination_matrix.py b/tests/parity/test_gam_pipeline_combination_matrix.py index c19838d4..3b2022b9 100644 --- a/tests/parity/test_gam_pipeline_combination_matrix.py +++ b/tests/parity/test_gam_pipeline_combination_matrix.py @@ -200,7 +200,7 @@ def test_stage_1_supported_numeric_factor_interactions_rebuild_on_newdata_like_m np.testing.assert_allclose(actual_se, np.asarray(expected["se"]).ravel(), atol=2e-8, rtol=2e-8) -@pytest.mark.parametrize("basis", ["cr", "cs", "cc", "ps", "tp", "ts"]) +@pytest.mark.parametrize("basis", ["cr", "cs", "cc", "cp", "ps", "tp", "ts"]) def test_stage_2_univariate_runtime_boundary_predictions_and_se_match_mgcv(basis): """Every supported univariate runtime matches behavior at and beyond fit bounds.""" data = _formula_data(seed=903, n=90)[["y", "x0"]].rename(columns={"x0": "x"}) diff --git a/tests/parity/test_mgcv_cp_combinations_parity.py b/tests/parity/test_mgcv_cp_combinations_parity.py new file mode 100644 index 00000000..e9774be1 --- /dev/null +++ b/tests/parity/test_mgcv_cp_combinations_parity.py @@ -0,0 +1,199 @@ +"""Integrated parity coverage for cyclic P-splines (``bs='cp'``).""" + +from __future__ import annotations + +import numpy as np +import pandas as pd +import pytest + +from nampy.gam import GAM +from nampy.gam.splines.univariate.ps import build_pspline_term_setup +from tests.mgcv_parity_utils import ( + _assert_basic_mgcv_parity, + _fit_nampy_model, + _fit_nampy_snapshot, + _run_mgcv_snapshot, +) + + +def _cp_data(seed=231, n=190): + rng = np.random.default_rng(seed) + x0 = rng.uniform(0.0, 2.0 * np.pi, size=n) + x1 = rng.uniform(-1.0, 5.0, size=n) + z = 0.7 + rng.uniform(-0.5, 0.8, size=n) + f = np.asarray(["a", "b", "c"], dtype=object)[np.arange(n) % 3] + f1 = np.asarray(["u", "v"], dtype=object)[np.arange(n) % 2] + y = ( + 0.3 + + z * np.sin(x0) + + 0.45 * np.cos(x1) + + 0.2 * (f == "b") + - 0.15 * (f1 == "v") + + rng.normal(scale=0.12, size=n) + ) + return pd.DataFrame({"y": y, "x0": x0, "x1": x1, "z": z, "f": f, "f1": f1}) + + +def _assert_snapshot_fit(actual, expected, *, atol=3e-8): + for key in ("response", "link"): + np.testing.assert_allclose( + actual["predictions"][key], expected["predictions"][key], atol=atol, rtol=atol + ) + np.testing.assert_allclose( + actual["fit"]["edf_total"], expected["fit"]["edf_total"], atol=atol, rtol=atol + ) + + +def test_cp_numeric_by_select_true_matches_mgcv(): + data = _cp_data(seed=232) + formula = 'y ~ s(x0, by=z, bs="cp", k=9)' + actual = _fit_nampy_snapshot(data, formula, "gaussian", "REML", select=True) + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML", select=True) + assert len(actual["fit"]["smoothing_params"]) == 2 + _assert_basic_mgcv_parity( + actual, + expected, + pred_atol=2e-8, + pred_rtol=2e-8, + sp_log_atol=2e-7, + criterion_atol=2e-8, + ) + + +def test_cp_centered_select_true_keeps_one_penalty_like_mgcv(): + data = _cp_data(seed=240) + formula = 'y ~ s(x0, bs="cp", k=9, sp=0.7)' + actual = _fit_nampy_snapshot( + data, formula, "gaussian", "fixed", select=True + ) + expected = _run_mgcv_snapshot( + data, formula, "gaussian", "REML", select=True + ) + assert len(actual["fit"]["smoothing_params"]) == 1 + _assert_snapshot_fit(actual, expected) + + +def test_cp_factor_by_fixed_sp_matches_mgcv(): + data = _cp_data(seed=233) + formula = 'y ~ s(x0, by=f, bs="cp", k=8, sp=0.7)' + actual = _fit_nampy_snapshot(data, formula, "gaussian", "fixed") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + _assert_snapshot_fit(actual, expected) + + +def test_cp_linked_terms_pool_basis_and_share_sp_like_mgcv(): + data = _cp_data(seed=234) + formula = ( + 'y ~ s(x0, bs="cp", k=9, id="periodic")' + ' + s(x1, bs="cp", k=9, id="periodic")' + ) + model = _fit_nampy_model(data, formula, "gaussian", "REML") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + assert len(model.smoothing_params) == 1 + runtimes = [ + term.predict_fn.__self__ + for term in model.gam_result_.compiled_model.compiled_terms + if term.term_type == "smooth" + ] + assert len(runtimes) == 2 + np.testing.assert_allclose(runtimes[0]._setup.knots, runtimes[1]._setup.knots) + pooled = np.concatenate([data["x0"].to_numpy(), data["x1"].to_numpy()]) + assert np.min(runtimes[0]._setup.knots) == pytest.approx(np.min(pooled)) + assert np.max(runtimes[0]._setup.knots) == pytest.approx(np.max(pooled)) + actual = model.parity_snapshot(X=data, include_covariances=True) + _assert_basic_mgcv_parity( + actual, + expected, + pred_atol=2e-8, + pred_rtol=2e-8, + sp_log_atol=2e-7, + criterion_atol=2e-8, + ) + + +def test_cp_fixed_term_has_no_penalty_and_matches_mgcv(): + data = _cp_data(seed=235) + formula = 'y ~ s(x0, bs="cp", k=8, fx=True)' + model = _fit_nampy_model(data, formula, "gaussian", "fixed") + assert model.gam_result_.compiled_model.compiled_penalties == () + actual = model.parity_snapshot(X=data, include_covariances=True) + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + _assert_snapshot_fit(actual, expected) + + +@pytest.mark.parametrize( + "formula", + [ + 'y ~ te(x0, x1, bs=["cp","cp"], k=[5,6], m=[1,2], sp=[0.6,0.8])', + 'y ~ ti(x0, x1, bs=["cp","cp"], k=[5,6], m=[1,2], sp=[0.6,0.8])', + ], + ids=["te", "ti"], +) +def test_cp_tensor_marginal_fixed_sp_fit_matches_mgcv(formula): + data = _cp_data(seed=236, n=170) + actual = _fit_nampy_snapshot(data, formula, "gaussian", "fixed") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + _assert_snapshot_fit(actual, expected, atol=2e-7) + + +@pytest.mark.parametrize( + "formula", + [ + 'y ~ s(f, x0, bs="fs", k=6, xt="cp", sp=[0.7,0.9])', + 'y ~ s(f, f1, x0, bs="sz", k=6, xt="cp", id="shared", sp=0.7)', + ], + ids=["fs", "sz"], +) +def test_cp_factor_smooth_base_fixed_sp_fit_matches_mgcv(formula): + data = _cp_data(seed=237, n=180) + actual = _fit_nampy_snapshot(data, formula, "gaussian", "fixed") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + _assert_snapshot_fit(actual, expected, atol=3e-7) + + +def test_cp_array_api_and_persistence_preserve_wrapped_prediction(tmp_path): + data = _cp_data(seed=238, n=150) + features = data[["x0", "x1"]] + model = GAM( + family="gaussian", + basis="cp", + k=8, + optimize_smoothing=False, + smoothing_params=[0.5, 0.8], + ).fit(X=features, y=data["y"].to_numpy(dtype=np.float64)) + newdata = pd.DataFrame({"x0": [-7.0, 0.0, 7.0], "x1": [-8.0, 1.0, 9.0]}) + expected = model.predict(newdata, type="link") + path = tmp_path / "cp.pkl" + model.save_model(path) + restored = GAM.load_model(path) + np.testing.assert_allclose(restored.predict(newdata, type="link"), expected) + + +@pytest.mark.parametrize( + "formula,message", + [ + ('y ~ s(x0, bs="cp", k=1, m=[2,1])', "basis dimension too small"), + ('y ~ s(x0, bs="cp", k=3, m=[2,3])', "penalty order too high"), + ], +) +def test_cp_invalid_orders_fail_loudly(formula, message): + data = _cp_data(seed=239, n=40) + with pytest.raises(ValueError, match=message): + GAM(formula=formula).fit(data=data) + + +def test_cp_knot_validation_and_uninformed_coefficient_warning(): + x = np.linspace(0.0, 1.0, 30) + kwargs = { + "feature_index": 0, + "feature_name": "x", + "bs_dim": 8, + "m": (2, 2), + "basis": "cp", + } + with pytest.raises(ValueError, match="knot range does not include data"): + build_pspline_term_setup(x, knots=[0.2, 0.8], **kwargs) + with pytest.raises(ValueError, match="there should be 9 supplied knots"): + build_pspline_term_setup(x, knots=np.linspace(0.0, 1.0, 8), **kwargs) + with pytest.warns(UserWarning, match="no.*information"): + build_pspline_term_setup(x, knots=[-100.0, 100.0], **kwargs) diff --git a/tests/reference_fixtures/mgcv/0102343adaf8b804e6ee2e7c0ab483a5634a4d58288dfdff2e9073ede53eec53.json.gz b/tests/reference_fixtures/mgcv/0102343adaf8b804e6ee2e7c0ab483a5634a4d58288dfdff2e9073ede53eec53.json.gz new file mode 100644 index 00000000..b4915550 Binary files /dev/null and b/tests/reference_fixtures/mgcv/0102343adaf8b804e6ee2e7c0ab483a5634a4d58288dfdff2e9073ede53eec53.json.gz differ diff --git a/tests/reference_fixtures/mgcv/06e67d7db6211f8efa71d8476faf5f36ba9bb5236fadf7eed1241100348cc5b5.json.gz b/tests/reference_fixtures/mgcv/06e67d7db6211f8efa71d8476faf5f36ba9bb5236fadf7eed1241100348cc5b5.json.gz new file mode 100644 index 00000000..d031922b Binary files /dev/null and b/tests/reference_fixtures/mgcv/06e67d7db6211f8efa71d8476faf5f36ba9bb5236fadf7eed1241100348cc5b5.json.gz differ diff --git a/tests/reference_fixtures/mgcv/0be540cc78c228c41fb4027e4a351b836b5385b50efbefc2c7292648f259934f.json.gz b/tests/reference_fixtures/mgcv/0be540cc78c228c41fb4027e4a351b836b5385b50efbefc2c7292648f259934f.json.gz new file mode 100644 index 00000000..b1c14ea1 Binary files /dev/null and b/tests/reference_fixtures/mgcv/0be540cc78c228c41fb4027e4a351b836b5385b50efbefc2c7292648f259934f.json.gz differ diff --git a/tests/reference_fixtures/mgcv/0ef818f5c0c0645f3582af3c8800c08f0f22bfcfb2a02de5793d702af681b39e.json.gz b/tests/reference_fixtures/mgcv/0ef818f5c0c0645f3582af3c8800c08f0f22bfcfb2a02de5793d702af681b39e.json.gz new file mode 100644 index 00000000..7db81832 Binary files /dev/null and b/tests/reference_fixtures/mgcv/0ef818f5c0c0645f3582af3c8800c08f0f22bfcfb2a02de5793d702af681b39e.json.gz differ diff --git a/tests/reference_fixtures/mgcv/1cf66a4576c9498ba2850d460d42b44d8eeaf7fd1da67f776f5f058c20a396d3.json.gz b/tests/reference_fixtures/mgcv/1cf66a4576c9498ba2850d460d42b44d8eeaf7fd1da67f776f5f058c20a396d3.json.gz new file mode 100644 index 00000000..f9eb9f5d Binary files /dev/null and b/tests/reference_fixtures/mgcv/1cf66a4576c9498ba2850d460d42b44d8eeaf7fd1da67f776f5f058c20a396d3.json.gz differ diff --git a/tests/reference_fixtures/mgcv/2ebce96ad51bc2308c2df3e2e635e758464ed97238696cfcf1b9c22c4f6331b9.json.gz b/tests/reference_fixtures/mgcv/2ebce96ad51bc2308c2df3e2e635e758464ed97238696cfcf1b9c22c4f6331b9.json.gz new file mode 100644 index 00000000..556070b7 Binary files /dev/null and b/tests/reference_fixtures/mgcv/2ebce96ad51bc2308c2df3e2e635e758464ed97238696cfcf1b9c22c4f6331b9.json.gz differ diff --git a/tests/reference_fixtures/mgcv/4370cfaddf37e346b4e6ef9083bba5c46be8282be51a3338f1cad80bb26e455d.json.gz b/tests/reference_fixtures/mgcv/4370cfaddf37e346b4e6ef9083bba5c46be8282be51a3338f1cad80bb26e455d.json.gz new file mode 100644 index 00000000..e9cd3d17 Binary files /dev/null and b/tests/reference_fixtures/mgcv/4370cfaddf37e346b4e6ef9083bba5c46be8282be51a3338f1cad80bb26e455d.json.gz differ diff --git a/tests/reference_fixtures/mgcv/4a39dbd5fb9848784b47a35cdf3bb5a6aa638d4c8db07296f74887c7528c205d.json.gz b/tests/reference_fixtures/mgcv/4a39dbd5fb9848784b47a35cdf3bb5a6aa638d4c8db07296f74887c7528c205d.json.gz new file mode 100644 index 00000000..2b06b87a Binary files /dev/null and b/tests/reference_fixtures/mgcv/4a39dbd5fb9848784b47a35cdf3bb5a6aa638d4c8db07296f74887c7528c205d.json.gz differ diff --git a/tests/reference_fixtures/mgcv/4c2d78a005ab3e06fadba7371e8203b8439d10e87adcc9252fe2565154ad969b.json.gz b/tests/reference_fixtures/mgcv/4c2d78a005ab3e06fadba7371e8203b8439d10e87adcc9252fe2565154ad969b.json.gz new file mode 100644 index 00000000..7763c164 Binary files /dev/null and b/tests/reference_fixtures/mgcv/4c2d78a005ab3e06fadba7371e8203b8439d10e87adcc9252fe2565154ad969b.json.gz differ diff --git a/tests/reference_fixtures/mgcv/5bdefc4353235bf4354c3dd4845e2d4d92c8a9352b53a62b64acef01269bfe90.json.gz b/tests/reference_fixtures/mgcv/5bdefc4353235bf4354c3dd4845e2d4d92c8a9352b53a62b64acef01269bfe90.json.gz new file mode 100644 index 00000000..5d84f6e6 Binary files /dev/null and b/tests/reference_fixtures/mgcv/5bdefc4353235bf4354c3dd4845e2d4d92c8a9352b53a62b64acef01269bfe90.json.gz differ diff --git a/tests/reference_fixtures/mgcv/6743c3eb0cf510edf14a540005184260ae4a2d71ade7a4c75887f20c24d8e354.json.gz b/tests/reference_fixtures/mgcv/6743c3eb0cf510edf14a540005184260ae4a2d71ade7a4c75887f20c24d8e354.json.gz new file mode 100644 index 00000000..1042e4b1 Binary files /dev/null and b/tests/reference_fixtures/mgcv/6743c3eb0cf510edf14a540005184260ae4a2d71ade7a4c75887f20c24d8e354.json.gz differ diff --git a/tests/reference_fixtures/mgcv/7033f1535cd827b645a593cc6db0d8046ee67705f9c70b7daa543cbff094ff98.json.gz b/tests/reference_fixtures/mgcv/7033f1535cd827b645a593cc6db0d8046ee67705f9c70b7daa543cbff094ff98.json.gz new file mode 100644 index 00000000..78a72451 Binary files /dev/null and b/tests/reference_fixtures/mgcv/7033f1535cd827b645a593cc6db0d8046ee67705f9c70b7daa543cbff094ff98.json.gz differ diff --git a/tests/reference_fixtures/mgcv/7c257db9cbb304215926619dffe5c7a636d51c64758cbf6b13279e86d6ae729f.json.gz b/tests/reference_fixtures/mgcv/7c257db9cbb304215926619dffe5c7a636d51c64758cbf6b13279e86d6ae729f.json.gz new file mode 100644 index 00000000..59decc70 Binary files /dev/null and b/tests/reference_fixtures/mgcv/7c257db9cbb304215926619dffe5c7a636d51c64758cbf6b13279e86d6ae729f.json.gz differ diff --git a/tests/reference_fixtures/mgcv/81c2643597263549fc7da7d9430f6214438b68113065b5eb6a73166fd3d3bbd7.json.gz b/tests/reference_fixtures/mgcv/81c2643597263549fc7da7d9430f6214438b68113065b5eb6a73166fd3d3bbd7.json.gz new file mode 100644 index 00000000..b5498a2f Binary files /dev/null and b/tests/reference_fixtures/mgcv/81c2643597263549fc7da7d9430f6214438b68113065b5eb6a73166fd3d3bbd7.json.gz differ diff --git a/tests/reference_fixtures/mgcv/8c9615e7f8cda28a977fc277e9d5f2a0255e48c2469726fd2f96c9cf8e137c2f.json.gz b/tests/reference_fixtures/mgcv/8c9615e7f8cda28a977fc277e9d5f2a0255e48c2469726fd2f96c9cf8e137c2f.json.gz new file mode 100644 index 00000000..26b37c12 Binary files /dev/null and b/tests/reference_fixtures/mgcv/8c9615e7f8cda28a977fc277e9d5f2a0255e48c2469726fd2f96c9cf8e137c2f.json.gz differ diff --git a/tests/reference_fixtures/mgcv/9a01c02034f7950fe93579e86db46aac42350946791f944f3185e76cb507c951.json.gz b/tests/reference_fixtures/mgcv/9a01c02034f7950fe93579e86db46aac42350946791f944f3185e76cb507c951.json.gz new file mode 100644 index 00000000..96f9312b Binary files /dev/null and b/tests/reference_fixtures/mgcv/9a01c02034f7950fe93579e86db46aac42350946791f944f3185e76cb507c951.json.gz differ diff --git a/tests/reference_fixtures/mgcv/9cc16601a52ecbecd7a045df5f7e1c02361167de3e8f78f96eb96b9757d92fec.json.gz b/tests/reference_fixtures/mgcv/9cc16601a52ecbecd7a045df5f7e1c02361167de3e8f78f96eb96b9757d92fec.json.gz new file mode 100644 index 00000000..651ca91d Binary files /dev/null and b/tests/reference_fixtures/mgcv/9cc16601a52ecbecd7a045df5f7e1c02361167de3e8f78f96eb96b9757d92fec.json.gz differ diff --git a/tests/reference_fixtures/mgcv/9ffe83ac5b4d70741d58cdd356d7c5d896a0f0e714997ca644b3dd3a48ae0c29.json.gz b/tests/reference_fixtures/mgcv/9ffe83ac5b4d70741d58cdd356d7c5d896a0f0e714997ca644b3dd3a48ae0c29.json.gz new file mode 100644 index 00000000..5a25841a Binary files /dev/null and b/tests/reference_fixtures/mgcv/9ffe83ac5b4d70741d58cdd356d7c5d896a0f0e714997ca644b3dd3a48ae0c29.json.gz differ diff --git a/tests/reference_fixtures/mgcv/bc062553471c566e3aa15008fb8c41216e51f496c1d3baede30bc10941060ecc.json.gz b/tests/reference_fixtures/mgcv/bc062553471c566e3aa15008fb8c41216e51f496c1d3baede30bc10941060ecc.json.gz new file mode 100644 index 00000000..d17c95e3 Binary files /dev/null and b/tests/reference_fixtures/mgcv/bc062553471c566e3aa15008fb8c41216e51f496c1d3baede30bc10941060ecc.json.gz differ diff --git a/tests/reference_fixtures/mgcv/c194c1779bee5c64ff8e221b5baa4d32e40040f3b809dba325ca93a3ccced397.json.gz b/tests/reference_fixtures/mgcv/c194c1779bee5c64ff8e221b5baa4d32e40040f3b809dba325ca93a3ccced397.json.gz new file mode 100644 index 00000000..9ca908fb Binary files /dev/null and b/tests/reference_fixtures/mgcv/c194c1779bee5c64ff8e221b5baa4d32e40040f3b809dba325ca93a3ccced397.json.gz differ diff --git a/tests/reference_fixtures/mgcv/c29b1b99a9f0670ec21e32d6dc067ce58364eae67901f615447a416ac59279a5.json.gz b/tests/reference_fixtures/mgcv/c29b1b99a9f0670ec21e32d6dc067ce58364eae67901f615447a416ac59279a5.json.gz new file mode 100644 index 00000000..395b79b2 Binary files /dev/null and b/tests/reference_fixtures/mgcv/c29b1b99a9f0670ec21e32d6dc067ce58364eae67901f615447a416ac59279a5.json.gz differ diff --git a/tests/reference_fixtures/mgcv/c3fc3837ed68e3275328a9c0e649fc19095efb056fb132e10e23bfbef4402e29.json.gz b/tests/reference_fixtures/mgcv/c3fc3837ed68e3275328a9c0e649fc19095efb056fb132e10e23bfbef4402e29.json.gz new file mode 100644 index 00000000..cf3ca233 Binary files /dev/null and b/tests/reference_fixtures/mgcv/c3fc3837ed68e3275328a9c0e649fc19095efb056fb132e10e23bfbef4402e29.json.gz differ diff --git a/tests/reference_fixtures/mgcv/d14d4e86c0cffe86fea2682bb119f8ce6d749509f2fffb379f7707c54be74f0f.json.gz b/tests/reference_fixtures/mgcv/d14d4e86c0cffe86fea2682bb119f8ce6d749509f2fffb379f7707c54be74f0f.json.gz new file mode 100644 index 00000000..264a6e39 Binary files /dev/null and b/tests/reference_fixtures/mgcv/d14d4e86c0cffe86fea2682bb119f8ce6d749509f2fffb379f7707c54be74f0f.json.gz differ diff --git a/tests/reference_fixtures/mgcv/d9c76231b23f5a6aec2a1197ba640c2e7ba3fb2b870d9bbbd49fda2371a55760.json.gz b/tests/reference_fixtures/mgcv/d9c76231b23f5a6aec2a1197ba640c2e7ba3fb2b870d9bbbd49fda2371a55760.json.gz new file mode 100644 index 00000000..65d90a4e Binary files /dev/null and b/tests/reference_fixtures/mgcv/d9c76231b23f5a6aec2a1197ba640c2e7ba3fb2b870d9bbbd49fda2371a55760.json.gz differ diff --git a/tests/reference_fixtures/mgcv/de589206367dc3f4baa0c730e5349aead0a1396da8df1d6769539d66287a3864.json.gz b/tests/reference_fixtures/mgcv/de589206367dc3f4baa0c730e5349aead0a1396da8df1d6769539d66287a3864.json.gz new file mode 100644 index 00000000..bb6891ca Binary files /dev/null and b/tests/reference_fixtures/mgcv/de589206367dc3f4baa0c730e5349aead0a1396da8df1d6769539d66287a3864.json.gz differ diff --git a/tests/reference_fixtures/mgcv/ee82f2b4c6bcb4f6b2698d4c95e015733e3df26b91ef03dde3b46d9387035d00.json.gz b/tests/reference_fixtures/mgcv/ee82f2b4c6bcb4f6b2698d4c95e015733e3df26b91ef03dde3b46d9387035d00.json.gz new file mode 100644 index 00000000..849eafbc Binary files /dev/null and b/tests/reference_fixtures/mgcv/ee82f2b4c6bcb4f6b2698d4c95e015733e3df26b91ef03dde3b46d9387035d00.json.gz differ diff --git a/tests/reference_fixtures/mgcv/fb980fbdcd9bb027f712cdfbb492b06316d08fd324abae11777c9a6706549dad.json.gz b/tests/reference_fixtures/mgcv/fb980fbdcd9bb027f712cdfbb492b06316d08fd324abae11777c9a6706549dad.json.gz new file mode 100644 index 00000000..fef6cc4e Binary files /dev/null and b/tests/reference_fixtures/mgcv/fb980fbdcd9bb027f712cdfbb492b06316d08fd324abae11777c9a6706549dad.json.gz differ diff --git a/tests/reference_fixtures/mgcv/ff8789f57b264e543f9bf3b37c2210afec728fba039cda5ccae2fba1c2f29873.json.gz b/tests/reference_fixtures/mgcv/ff8789f57b264e543f9bf3b37c2210afec728fba039cda5ccae2fba1c2f29873.json.gz new file mode 100644 index 00000000..4a8ae3f8 Binary files /dev/null and b/tests/reference_fixtures/mgcv/ff8789f57b264e543f9bf3b37c2210afec728fba039cda5ccae2fba1c2f29873.json.gz differ diff --git a/tests/scam/test_generic_transform_contracts.py b/tests/scam/test_generic_transform_contracts.py index ed05cd5b..4666ca0c 100644 --- a/tests/scam/test_generic_transform_contracts.py +++ b/tests/scam/test_generic_transform_contracts.py @@ -137,11 +137,12 @@ def test_ar1_reml_rejects_missing_correlation_likelihood_terms(): ).fit(data=data) -def test_ordinary_pspline_exposes_exact_derivative_provider_at_new_data(): +@pytest.mark.parametrize("basis", ["ps", "cp"]) +def test_pspline_exposes_exact_derivative_provider_at_new_data(basis): x = np.linspace(-1.5, 2.0, 60) data = pd.DataFrame({"y": np.sin(x), "x": x}) model = GAM( - formula='y ~ s(x, bs="ps", k=9, m=c(2, 2))', + formula=f'y ~ s(x, bs="{basis}", k=9, m=c(2, 2))', family="gaussian", smoothing_params=[0.8], ).fit(data=data) @@ -164,7 +165,7 @@ def test_ordinary_pspline_exposes_exact_derivative_provider_at_new_data(): assert derivative.derivative_matrix.shape[0] == len(new_data) -@pytest.mark.parametrize("basis", ["ps", "cr", "cc"]) +@pytest.mark.parametrize("basis", ["ps", "cp", "cr", "cc"]) def test_linear_functional_smooth_is_available_through_generic_by_contract(basis): rng = np.random.default_rng(91) locations = np.tile(np.linspace(-1.0, 1.0, 11), (28, 1)) diff --git a/tests/smooths/test_mgcv_raw_constructor_parity.py b/tests/smooths/test_mgcv_raw_constructor_parity.py index b1f2f682..899ddc5b 100644 --- a/tests/smooths/test_mgcv_raw_constructor_parity.py +++ b/tests/smooths/test_mgcv_raw_constructor_parity.py @@ -38,6 +38,7 @@ cyclic_cubic_predict_matrix, ) from nampy.gam.splines.univariate.ps import ( + cyclic_pspline_knots, pspline_knots, ) from tests.mgcv_invariant_policy import canonicalize_raw_representation_state @@ -152,6 +153,33 @@ def _build(data): return _build +def _cyclic_pspline_supplied_knots(column: str, bs_dim: int, basis_order: int): + def _build(data): + vals = np.asarray(data[column], dtype=np.float64) + return { + str(column): cyclic_pspline_knots( + vals, + bs_dim=int(bs_dim), + basis_order=int(basis_order), + supplied_knots=None, + ) + } + + return _build + + +def _cyclic_irregular_unsorted_knots(column: str): + def _build(_data): + return { + str(column): np.asarray( + [6.4, 0.0, 0.15, 0.6, 1.4, 2.1, 3.8, 4.6, 5.9], + dtype=np.float64, + ) + } + + return _build + + def _paired_feature_knots(columns, n_knots: int): cols = [str(col) for col in columns] @@ -334,6 +362,56 @@ def _build_ps_case_matrix(): ] +def _build_cp_case_matrix(): + return [ + _case( + "cp_default_k_default_m", + _factory(_make_cyclic_data, seed=133), + 'y ~ s(x, bs="cp")', + ), + _case( + "cp_m_scalar_1", + _factory(_make_cyclic_data, seed=134), + 'y ~ s(x, bs="cp", k=9, m=1)', + ), + _case( + "cp_m_ordered", + _factory(_make_cyclic_data, seed=135), + 'y ~ s(x, bs="cp", k=10, m=[2, 3])', + ), + _case( + "cp_zero_order_penalty_metadata", + _factory(_make_cyclic_data, seed=138), + 'y ~ s(x, bs="cp", k=7, m=[2, 0])', + ), + _case( + "cp_extra_m_entries_retained", + _factory(_make_cyclic_data, seed=139), + 'y ~ s(x, bs="cp", k=8, m=[2, 1, 9])', + ), + _case( + "cp_endpoint_knots", + _factory(_make_cyclic_data, seed=136), + 'y ~ s(x, bs="cp", k=8, m=[2, 2])', + knots_factory=_cyclic_endpoint_knots("x"), + ), + _case( + "cp_full_knots", + _factory(_make_cyclic_data, seed=137), + 'y ~ s(x, bs="cp", k=8, m=[2, 2])', + knots_factory=_cyclic_pspline_supplied_knots( + "x", bs_dim=8, basis_order=2 + ), + ), + _case( + "cp_irregular_unsorted_full_knots", + _factory(_make_cyclic_data, seed=140), + 'y ~ s(x, bs="cp", k=8, m=[2, 2])', + knots_factory=_cyclic_irregular_unsorted_knots("x"), + ), + ] + + def _build_tprs_case_matrix(): cases = [] for basis, seed_base in [("tp", 60), ("ts", 80)]: @@ -523,6 +601,7 @@ def _build_factor_smooth_case_matrix(): ("cs", "cs"), ("cc", "cc"), ("ps", {"bs": "ps", "m": 2, "k": 7}), + ("cp", {"bs": "cp", "m": 2, "k": 7}), ("ts", "ts"), ] for label, xt_spec in fs_xt_cases: @@ -531,7 +610,7 @@ def _build_factor_smooth_case_matrix(): f"fs_base_{label}", _make_fs_data, f'y ~ s(f, x, bs="fs", xt={repr(xt_spec)})', - atol=1e-8 if label in {"ps", "ts"} else 1e-10, + atol=1e-8 if label in {"ps", "cp", "ts"} else 1e-10, ) ) cases.append( @@ -539,7 +618,7 @@ def _build_factor_smooth_case_matrix(): f"sz_base_{label}", _make_sz_data, f'y ~ s(f1, f2, x, bs="sz", k=6, xt={repr(xt_spec)})', - atol=1e-8 if label in {"ps", "ts"} else 1e-10, + atol=1e-8 if label in {"ps", "cp", "ts"} else 1e-10, ) ) @@ -601,6 +680,18 @@ def _build_tensor_case_matrix(): _factory(_make_gaussian_data, seed=803, n=90), 'y ~ te(x0, x1, bs=["ps", "ps"], k=[6, 6], m=[[2, 2], [2, 3]])', ), + _case( + "te_cp_cp_m", + _factory(_make_gaussian_data, seed=818, n=90), + 'y ~ te(x0, x1, bs=["cp", "cp"], k=[5, 6], m=[[2, 1], [2, 2]])', + atol=1e-8, + ), + _case( + "ti_cp_cp_m", + _factory(_make_gaussian_data, seed=819, n=90), + 'y ~ ti(x0, x1, bs=["cp", "cp"], k=[5, 6], m=[[2, 1], [2, 2]])', + atol=1e-8, + ), _case( "te_tp_ts_m", _factory(_make_gaussian_data, seed=804, n=90), @@ -664,6 +755,7 @@ def _build_tensor_case_matrix(): CASES = [ *_build_cubic_case_matrix(), *_build_ps_case_matrix(), + *_build_cp_case_matrix(), *_build_tprs_case_matrix(), *_build_re_case_matrix(), *_build_factor_smooth_case_matrix(), @@ -803,14 +895,22 @@ def _serialize_ps_raw(term): setup = term._setup B = np.asarray(setup.basis_train, dtype=np.float64) return _common_raw_state( - "pspline.smooth", + ( + "cpspline.smooth" + if str(setup.basis_name).lower() == "cp" + else "pspline.smooth" + ), B, [np.asarray(setup.penalty, dtype=np.float64)], rank=int(setup.rank), null_space_dim=int(B.shape[1] - setup.rank), extra={ "knots": np.asarray(setup.knots, dtype=np.float64), - "m": _scalar_or_list([int(setup.basis_order), int(setup.penalty_order)]), + "m": _scalar_or_list( + list(setup.orders) + if setup.orders + else [int(setup.basis_order), int(setup.penalty_order)] + ), }, ) diff --git a/tests/smooths/test_mgcv_smoothcon_parity.py b/tests/smooths/test_mgcv_smoothcon_parity.py index 340d4c37..864a29be 100644 --- a/tests/smooths/test_mgcv_smoothcon_parity.py +++ b/tests/smooths/test_mgcv_smoothcon_parity.py @@ -788,6 +788,80 @@ def test_gaussian_cc_reml_matches_mgcv(self): # --------------------------------------------------------------------------- +class TestCyclicPSplineSmooth: + """Cyclic P-spline (bs='cp') constructor and fit parity.""" + + @staticmethod + def _make_data(seed=181, n=190): + rng = np.random.default_rng(seed) + x = rng.uniform(0.0, 2.0 * np.pi, size=n) + y = np.sin(x) + 0.25 * np.cos(2.0 * x) + rng.normal(scale=0.12, size=n) + return pd.DataFrame({"y": y, "x": x}) + + def test_cp_smoothcon_basis_and_penalty_match_mgcv(self): + data = self._make_data() + formula = 'y ~ s(x, bs="cp", k=11, m=[2,1])' + smooth_expr = 's(x, bs="cp", k=11, m=c(2,1))' + design = _compile_formula_design(data, formula) + expected_basis = _run_mgcv_smoothcon_matrix(data, smooth_expr) + expected_penalties = _run_mgcv_smoothcon_penalties( + data, smooth_expr, absorb_cons=True, scale_penalty=True + ) + + np.testing.assert_allclose( + design.design_matrix, expected_basis["X"], atol=1e-10, rtol=1e-10 + ) + actual = [pb.matrix for pb in design.compiled_penalties] + assert len(actual) == len(expected_penalties["S"]) == 1 + np.testing.assert_allclose( + actual[0], expected_penalties["S"][0], atol=1e-10, rtol=1e-10 + ) + + def test_cp_point_constraint_matches_mgcv(self): + data = self._make_data(seed=182) + formula = 'y ~ s(x, bs="cp", k=9, pc=0.0, sp=0.6)' + actual = _fit_nampy_snapshot(data, formula, "gaussian", "fixed") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + np.testing.assert_allclose( + actual["predictions"]["response"], + expected["predictions"]["response"], + atol=2e-10, + rtol=2e-10, + ) + + def test_cp_fixed_sp_fit_matches_mgcv(self): + data = self._make_data(seed=183) + formula = 'y ~ s(x, bs="cp", k=10, m=[2,2], sp=0.7)' + actual = _fit_nampy_snapshot(data, formula, "gaussian", "fixed") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + np.testing.assert_allclose( + actual["predictions"]["response"], + expected["predictions"]["response"], + atol=2e-10, + rtol=2e-10, + ) + np.testing.assert_allclose( + actual["fit"]["cov_bayes"], + expected["fit"]["cov_bayes"], + atol=2e-10, + rtol=2e-10, + ) + + def test_cp_reml_fit_matches_mgcv(self): + data = self._make_data(seed=184, n=210) + formula = 'y ~ s(x, bs="cp", k=11, m=[2,2])' + actual = _fit_nampy_snapshot(data, formula, "gaussian", "REML") + expected = _run_mgcv_snapshot(data, formula, "gaussian", "REML") + _assert_basic_mgcv_parity( + actual, + expected, + pred_atol=2e-9, + pred_rtol=2e-9, + sp_log_atol=2e-8, + criterion_atol=2e-8, + ) + + class TestPSplineSmooth(_SharedTestPSplineSmooth): """P-spline (bs='ps') standalone parity against mgcv."""