Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions docs/representations/spline_expansions.md
Original file line number Diff line number Diff line change
Expand Up @@ -23,8 +23,8 @@ knot positions by `placement_strategy` (see
For every univariate spline in this section, `output_dim` is the number of output columns
**per input feature**, not the total. A `(n_samples, 3)` input produces
`(n_samples, 3 * output_dim)` output (plus one extra column per feature if
`include_bias=True`). Feature names are suffixed per input, for example `x0_bs0, x0_bs1, ...`
for a B-spline on column `x0`.
`include_bias=True`). Feature names are suffixed per input, for example `age_bs0, age_bs1, ...`
for a B-spline fitted on a DataFrame column `age` (`x0_bs0, x0_bs1, ...` for array input).
```

## B-spline
Expand Down
58 changes: 54 additions & 4 deletions pretab/compose/inspection.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,21 +11,25 @@
from typing import Any

import numpy as np
import pandas as pd
from sklearn.pipeline import FeatureUnion, Pipeline

from ..core.logging import get_logger
from ..core.representation import FeatureLineage
from .registry import TRANSFORMER_REGISTRY

logger = get_logger(__name__)

__all__ = [
"block_name",
"block_uses_target",
"build_feature_info",
"build_feature_lineage",
"build_transformer_summary",
"clean_feature_names",
"feature_names_out",
"get_output_slices",
"representation_leaf",
]


Expand Down Expand Up @@ -177,6 +181,32 @@ def _separate_state_branches(transformer):
return branches["representation"], branches["missing"]


def representation_leaf(transformer):
"""Return the last step of a fitted per-column block's representation.

For a separate-state / missing-indicator union this is the last step of its
``"representation"`` pipeline (the raw missing indicator beside it is not
part of the representation); for a pipeline it is the last step, and any
other block is returned unchanged.
"""
separate_state = _separate_state_branches(transformer)
representation = transformer if separate_state is None else separate_state[0]
return representation.steps[-1][1] if hasattr(representation, "steps") else representation


def _probe_input(step, value):
"""Return a one-row probe input for measuring a fitted step's output width.

A step fitted directly on the ColumnTransformer's DataFrame column records
``feature_names_in_`` and, like scikit-learn, warns on input without feature
names, so its probe is a DataFrame carrying the fitted names.
"""
names = getattr(step, "feature_names_in_", None)
if names is None:
return np.full((1, 1), value)
return pd.DataFrame(np.full((1, len(names)), value), columns=names)


def build_feature_info(column_transformer, *, embeddings, embedding_dimensions):
"""Collect per-feature metadata (preprocessing, dimension, categories).

Expand Down Expand Up @@ -230,7 +260,7 @@ def build_feature_info(column_transformer, *, embeddings, embedding_dimensions):
):
last_step = representation_pipeline.steps[-1][1]
if hasattr(last_step, "transform"):
dummy_input = np.zeros((1, 1)) + 1e-05
dummy_input = _probe_input(last_step, 1e-05)
try:
transformed_feature = last_step.transform(dummy_input)
dimension = transformed_feature.shape[1]
Expand Down Expand Up @@ -275,7 +305,7 @@ def build_feature_info(column_transformer, *, embeddings, embedding_dimensions):
else:
last_step = representation_pipeline.steps[-1][1]
if hasattr(last_step, "transform"):
dummy_input = np.zeros((1, 1))
dummy_input = _probe_input(last_step, 0.0)
try:
transformed_feature = last_step.transform(dummy_input)
dimension = transformed_feature.shape[1]
Expand Down Expand Up @@ -348,8 +378,9 @@ def _resolve_block_representation(pipeline, columns):
"""Return ``(family, component, uses_target, is_interaction)`` for a block.

The representation-bearing step is the last pipeline step exposing a
``get_representation_spec`` (a PreTab transformer) or a known scikit-learn
step name; helper steps such as imputers and float casts are skipped.
``get_representation_spec`` (a PreTab transformer), a known scikit-learn
step name, or a registered method that consumes the target; helper steps
such as imputers and float casts are skipped.
"""
steps = pipeline.steps if hasattr(pipeline, "steps") else [("_", pipeline)]
for step_name, transformer in reversed(steps):
Expand All @@ -359,9 +390,28 @@ def _resolve_block_representation(pipeline, columns):
if step_name in _STEP_FAMILY:
family, component = _STEP_FAMILY[step_name]
return family, component, False, False
registered = TRANSFORMER_REGISTRY.get(step_name)
if registered is not None and registered.target_usage != "forbidden":
# A registered class without a RepresentationSpec (e.g. scikit-learn's
# TargetEncoder): its declared supervision says whether it used y, so
# cross-fitting refits it per fold instead of reusing the all-data fit.
uses_target = registered.target_usage == "required" or bool(getattr(transformer, "target_aware", False))
component = "category" if registered.is_categorical else "basis"
return step_name, component, uses_target, registered.is_multivariate
return "passthrough", "raw", False, False


def block_uses_target(transformer, columns) -> bool:
"""Whether a fitted per-column block's representation consumed the target ``y``.

For a separate-state / missing-indicator union the representation branch
decides; helper steps such as imputers and scalers never use the target.
"""
separate_state = _separate_state_branches(transformer)
representation = separate_state[0] if separate_state is not None else transformer
return bool(_resolve_block_representation(representation, columns)[2])


def _passthrough_source(columns, offset, feature_names_in):
"""Resolve the source feature name for a passthrough / remainder column."""
column = columns[offset] if offset < len(columns) else columns[-1]
Expand Down
12 changes: 10 additions & 2 deletions pretab/compose/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,9 @@ class RepresentationSearchCV(BaseEstimator):
Scoring passed to :func:`sklearn.metrics.check_scoring`; ``None`` uses the
estimator's ``score`` method.
preprocessor_params : dict or None, default=None
Extra keyword arguments forwarded to every :class:`Preprocessor`.
Extra keyword arguments forwarded to every :class:`Preprocessor`. When
``estimator`` is a classifier, ``task`` defaults to ``"classification"`` so
target-aware placement treats ``y`` as class labels.
random_state : int or None, default=None
Seed forwarded to each :class:`Preprocessor`.

Expand All @@ -74,9 +76,15 @@ def __init__(self, estimator, methods, *, cv=5, scoring=None, preprocessor_param
self.random_state = random_state

def _make_preprocessor(self, method):
"""Build a Preprocessor for ``method`` with the shared parameters."""
"""Build a Preprocessor for ``method`` with the shared parameters.

Target-aware placement follows the downstream estimator: a classifier makes
it ``task="classification"`` unless ``preprocessor_params`` sets ``task``.
"""
params = dict(self.preprocessor_params or {})
params.setdefault("random_state", self.random_state)
if is_classifier(self.estimator):
params.setdefault("task", "classification")
return Preprocessor(numerical_method=method, **params)

def fit(self, X, y=None):
Expand Down
56 changes: 44 additions & 12 deletions pretab/compose/serialize.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

import dataclasses
import importlib
import math
from collections import UserList
from typing import Any, cast

Expand All @@ -27,8 +28,9 @@
from ..core.parameters import UNSET
from ..core.policy import RepresentationPolicy
from ..core.representation import FeatureLineage, RepresentationSpec
from ..exceptions import PretabError, PretabSerializationError
from ..exceptions import PretabSerializationError
from ..placement.base import PlacementResult
from .inspection import representation_leaf
from .registry import TransformerSpec

SCHEMA_VERSION = 1
Expand Down Expand Up @@ -88,14 +90,23 @@ def _resolve(dotted: str):


# --- encoding ------------------------------------------------------------
def _encode_float(value):
"""Return ``value``, tagging NaN / +-inf, which strict JSON cannot represent."""
if math.isfinite(value):
return value
return {"__float__": repr(float(value))}


def _encode(obj):
if isinstance(obj, np.bool_):
return bool(obj)
if isinstance(obj, np.integer):
return int(obj)
if isinstance(obj, np.floating):
return float(obj)
if obj is None or isinstance(obj, (bool, int, float, str)):
return _encode_float(float(obj))
if isinstance(obj, float):
return _encode_float(obj)
if obj is None or isinstance(obj, (bool, int, str)):
return obj
if obj is UNSET:
return {"__unset__": True}
Expand All @@ -105,6 +116,8 @@ def _encode(obj):
# tolist() keeps the elements of an object array as they are (e.g.
# numpy scalars such as np.bool_), so encode each one.
data = _map_elements(data, obj.ndim, _encode)
elif obj.dtype.kind == "f" and not np.isfinite(obj).all():
data = _map_elements(data, obj.ndim, _encode_float)
return {"__ndarray__": {"dtype": obj.dtype.str, "shape": list(obj.shape), "data": data}}
if isinstance(obj, np.dtype):
return {"__npdtype__": obj.str}
Expand Down Expand Up @@ -151,7 +164,9 @@ def _json_safe(obj):
Keeps JSON-native values verbatim and falls back to the tagged :func:`_encode`
form only for exotic values. The result is informational and never decoded.
"""
if obj is None or isinstance(obj, (bool, int, float, str)):
if isinstance(obj, float):
return _encode_float(obj)
if obj is None or isinstance(obj, (bool, int, str)):
return obj
if isinstance(obj, dict) and all(isinstance(k, str) for k in obj):
return {k: _json_safe(v) for k, v in obj.items()}
Expand All @@ -173,7 +188,10 @@ def _decode_ndarray(payload: dict) -> np.ndarray:
for index, element in enumerate(elements):
arr[index] = element
return arr.reshape(shape)
arr = np.array(payload["data"], dtype=dtype)
data = payload["data"]
if dtype.kind == "f":
data = _map_elements(data, len(payload["shape"]), _decode)
arr = np.array(data, dtype=dtype)
return arr.reshape(payload["shape"])


Expand All @@ -187,6 +205,8 @@ def _decode(obj):
return _decode_ndarray(obj["__ndarray__"])
if "__unset__" in obj:
return UNSET
if "__float__" in obj:
return float(obj["__float__"])
if "__npdtype__" in obj:
return np.dtype(obj["__npdtype__"])
if "__type__" in obj:
Expand Down Expand Up @@ -262,23 +282,35 @@ def _library_versions() -> dict:


def _representation_summary(preprocessor) -> list:
"""Best-effort declarative per-representation summary (family/columns/locations)."""
"""Declarative per-representation summary (family/columns/locations).

One entry per block whose representation ends in a PreTab transformer (for a
missing-indicator / separate-state block, its representation branch; see
:func:`~pretab.compose.inspection.representation_leaf`), built from its
:class:`RepresentationSpec` with the block's column names as input features
(as the feature lineage does). Blocks whose representation ends in a step
without ``get_representation_spec`` (scikit-learn scalers and encoders,
embeddings) have no entry. The steps are fitted, so a spec that fails to
build is a bug: its error propagates instead of silently dropping the entry.
"""
summary: list = []
column_transformer = getattr(preprocessor, "column_transformer_", None)
if column_transformer is None:
return summary
for name, transformer, columns in column_transformer.transformers_:
if name == "remainder":
continue
leaf = transformer.steps[-1][1] if hasattr(transformer, "steps") else transformer
leaf = representation_leaf(transformer)
spec_fn = getattr(leaf, "get_representation_spec", None)
if spec_fn is None:
continue
try:
entry = spec_fn().to_dict()
except (PretabError, ValueError, AttributeError, TypeError, KeyError):
continue
entry["columns"] = [str(col) for col in columns]
block_columns = [str(col) for col in columns]
# A leaf with more inputs than its block has columns (one fitted behind
# the imputer's built-in missing indicator, as before the #62 fix) names
# its own inputs.
fits_block = getattr(leaf, "n_features_in_", len(block_columns)) == len(block_columns)
entry = spec_fn(input_features=block_columns if fits_block else None).to_dict()
entry["columns"] = block_columns
summary.append(entry)
return summary

Expand Down
36 changes: 16 additions & 20 deletions pretab/core/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,12 +10,11 @@
from sklearn.base import BaseEstimator, TransformerMixin
from sklearn.utils.validation import check_is_fitted

from ..exceptions import invalid_param_error
from .adaptive import AdaptiveResolutionMixin
from .parameters import AliasResolverMixin
from .policy import RepresentationPolicy, apply_constant_policy
from .representation import RepresentationSpecMixin
from .validation import validate_2d_allow_nan
from .validation import resolve_input_features, validate_2d_allow_nan

__all__ = ["BasePreTabTransformer"]

Expand Down Expand Up @@ -58,19 +57,24 @@ class BasePreTabTransformer(
def _resolved_policy(self) -> RepresentationPolicy:
"""Return the effective edge-case policy for this instance.

An explicit ``policy`` constructor argument, on the transformers that expose
one, always wins verbatim over the class-level defaults below. Otherwise, the
shared default policy is narrowed by this family's ``_constant_policy`` /
``_out_of_range_policy`` class attributes, which record each family's
historical, non-configurable default behavior.
The shared default policy is narrowed by this family's ``_constant_policy``
/ ``_out_of_range_policy`` class attributes, which record each family's
historical default behavior. An explicit ``policy`` constructor argument, on
the transformers that expose one, then applies on top: a mapping overrides
only the axes it names (so ``{"constant": "error"}`` keeps the family's
out-of-range handling), and a :class:`RepresentationPolicy` instance is used
verbatim.
"""
instance_policy = getattr(self, "policy", None)
if instance_policy is not None:
return RepresentationPolicy.resolve(instance_policy)
return self._policy.merge(
family_policy = self._policy.merge(
constant=self._constant_policy,
out_of_range=self._out_of_range_policy,
)
instance_policy = getattr(self, "policy", None)
if instance_policy is None:
return family_policy
if isinstance(instance_policy, dict):
return family_policy.merge(**instance_policy)
return RepresentationPolicy.resolve(instance_policy)

def _validate(self, X, *, reset: bool):
"""Validate ``X`` through the shared NaN-aware validator."""
Expand All @@ -90,15 +94,7 @@ def _output_sizes(self) -> list[int]:
def get_feature_names_out(self, input_features=None):
"""Return output feature names of the form ``{feature}_{suffix}{j}``."""
check_is_fitted(self, "n_features_in_")
if input_features is None:
input_features = [f"x{i}" for i in range(self.n_features_in_)]
elif len(input_features) != self.n_features_in_:
raise invalid_param_error(
type(self).__name__,
"get_feature_names_out.input_features",
len(input_features),
f"must have exactly {self.n_features_in_} entries (one per input feature)",
)
input_features = resolve_input_features(self, input_features)
suffix = self._feature_suffix()
names = []
for feature, n_cols in zip(input_features, self._output_sizes(), strict=False):
Expand Down
18 changes: 18 additions & 0 deletions pretab/core/knots.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
__all__ = [
"basis_to_knots",
"bspline_basis",
"extrapolate_bspline_rows",
"generate_internal_knots",
"quantile_knots",
"select_knots",
Expand Down Expand Up @@ -58,6 +59,23 @@ def bspline_basis(x: np.ndarray, knots: np.ndarray, degree: int, i: int, last: i
return term1 + term2


def extrapolate_bspline_rows(basis: np.ndarray, x: np.ndarray, knots: np.ndarray, degree: int) -> np.ndarray:
"""Fill the rows of ``basis`` whose ``x`` lies outside the knot span by extrapolation.

:func:`bspline_basis` is zero outside ``[knots[0], knots[-1]]``. When a policy lets
out-of-range values through (``"extrapolate"`` / ``"warn"``), those rows are
evaluated instead by extending the boundary polynomial pieces, which keeps the
basis a partition of unity there. Rows inside the span are left untouched.
"""
outside = (x < knots[0]) | (x > knots[-1])
if outside.any():
from scipy.interpolate import BSpline

n_basis = len(knots) - degree - 1
basis[outside] = BSpline(knots, np.eye(n_basis), degree, extrapolate=True)(x[outside])
return basis


def basis_to_knots(n_basis: int, degree: int) -> int:
"""Number of internal knots implied by ``n_basis`` basis functions of ``degree``."""
return max(0, n_basis - degree - 1)
Expand Down
2 changes: 2 additions & 0 deletions pretab/core/policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
The defaults reproduce each family's historical behaviour (B/M/I-spline, P-spline, and
tensor-product clip out-of-range inputs; natural-cubic and cubic-regression extrapolate;
constant columns pass through everywhere), so leaving ``policy`` unset changes nothing.
A transformer's ``policy`` given as a mapping overrides only the axes it names on top of
those family defaults; a :class:`RepresentationPolicy` instance is used verbatim.
Transformers may narrow specific axes through class-level override attributes without
exposing a new constructor parameter (see :class:`~pretab.core.base.BasePreTabTransformer`).

Expand Down
Loading
Loading