Contributing¶
Development setup¶
Clone the repo and install in editable mode with dev dependencies:
Run the tests to check everything works:
Running the docs locally¶
Install the docs dependency and start the dev server:
This launches a local preview at http://127.0.0.1:8000/skeights/
that auto-reloads when you edit any docs files.
Adding support for a new estimator¶
skeights uses a handler pattern to support different estimator types. Each handler is a class that knows how to decompose a specific estimator into JSON state + numpy arrays, and how to put them back together.
The contract¶
Subclass skeights.EstimatorHandler and implement five methods:
from typing import Any
import numpy as np
from sklearn.base import BaseEstimator
from skeights._handler import EstimatorHandler
class MyHandler(EstimatorHandler):
def handles(self, estimator: BaseEstimator) -> bool:
"""Return True if this handler owns the given estimator type."""
return isinstance(estimator, MyEstimatorClass)
def collect_state(
self, estimator: BaseEstimator, prefix: str, format: str | None = None
) -> dict[str, Any]:
"""Extract JSON-safe scalar state (hyperparameters, metadata).
Returns a dict of scalar values that will be saved to JSON.
Use the prefix to namespace keys (important for pipelines).
"""
...
def restore_state(
self, estimator: BaseEstimator, fitted_state: dict[str, Any], prefix: str
) -> None:
"""Restore scalar state onto an estimator skeleton."""
...
def extract_arrays(
self, estimator: BaseEstimator, prefix: str, format: str | None = None
) -> dict[str, np.ndarray]:
"""Extract numpy arrays (weights, coefficients, etc).
Returns a dict of numpy arrays that will be saved to safetensors.
"""
...
def restore_arrays(
self,
estimator: BaseEstimator,
arrays: dict[str, np.ndarray],
prefix: str,
fitted_state: dict[str, Any] | None = None,
) -> None:
"""Restore numpy arrays onto an estimator."""
...
Registering the handler¶
Add your handler to the list in _core.py:
def _get_handlers() -> list[EstimatorHandler]:
from skeights._mymodule import MyHandler
# ... existing handlers ...
return [
# ... existing handlers ...
MyHandler(),
]
Key concepts¶
You only handle your estimator: pipelines, TransformedTargetRegressor,
and other composite estimators are handled by the core dispatch layer.
It walks the structure and calls each handler for the individual
estimator it owns. Your handler just needs to serialize and restore
a single estimator type. The composition is automatic.
Prefix: every key in the state dict and arrays dict is prefixed
with a path like scaler/ or model/. The core layer manages this
as it walks through composite estimators, giving each step its own
namespace. Always use f"{prefix}{key}" when building keys.
State vs arrays: anything that's a scalar, string, or small list goes in state (saved as JSON). Anything that's a numpy array goes in arrays (saved as safetensors). The split is important because JSON is human-readable and diffable, while safetensors is efficient for large numeric data.
Format: the format parameter lets handlers support multiple
serialization strategies. For example, LightGBM and XGBoost support
format="native" (library's own format) and the default columnar
tensors format.
Example: a simple handler¶
_mlp.py is the simplest handler to use as a reference. It
serializes MLP regressors and classifiers by:
- Saving scalar state (
n_layers_,n_outputs_, etc.) to JSON - Saving weight matrices (
coefs_,intercepts_) as numpy arrays
Testing¶
Add round-trip tests in tests/ that:
- Fit a model
- Save it with
skeights.save() - Load it with
skeights.load() - Check predictions match the original
See tests/test_mlp.py for a straightforward example.