Model wrappers¶
TestContext talks to models through a small adapter interface
(modeltest.wrappers):
predict(X) -> class labels
predict_proba(X) -> probability estimates (optional)
wrap(model) returns a normalized object implementing that interface,
dispatching on the framework. Out of the box it handles:
- scikit-learn estimators — any
BaseEstimatorwithpredict;predict_probais used when available (label probabilities viaSklearnClassifier). - sklearn
Pipelines — the same adapter works; the explainability tests additionally unwrap the pipeline so SHAP sees the final estimator over engineered features. - PyTorch
nn.Module—predictreturns argmax class indices,predict_probareturns softmax probabilities. The model is switched toeval()and gradients are disabled automatically. - Keras / TensorFlow —
predictthresholds single-output models at 0.5, or takes argmax for multiclass (multiclass=Trueforces it). - Anything else — the fallback assumes the model already exposes
predict/predict_proba(e.g. your own adapter), so custom model classes work without registration.
from modeltest import ModelSuite
from modeltest.scenarios import MinimumAccuracyTest
from modeltest.wrappers import wrap
model = wrap(my_custom_model) # optional — suites wrap automatically
suite = ModelSuite(name="suite", tests=[MinimumAccuracyTest()])
result = suite.run(model, X_val, y_val)
Tuning adapters¶
wrap forwards keyword arguments to the adapter:
| Adapter | kwarg | Purpose |
|---|---|---|
TorchModel |
input_key |
Key when the model expects a dict input. |
TorchModel |
device |
Device to move inputs to before forward. |
KerasModel |
multiclass |
Force argmax decoding even for 2 outputs. |
model = wrap(torch_model, device="cuda")
Feature-name filtering¶
Wrappers expose the model's feature_names_in_ when available. The context
uses it to filter X_val down to the exact features the model expects —
in the right order — before calling it. Extra helper columns (group columns,
IDs, timestamps) therefore coexist peacefully in your validation frame.
Writing your own adapter¶
Subclass ModelWrapper and implement predict (and optionally
predict_proba). Instances of ModelWrapper are passed through wrap
untouched, so you can pre-configure them and hand them to suite.run:
from modeltest.wrappers import ModelWrapper
class MyModelWrapper(ModelWrapper):
def predict(self, X):
return self.model.classify(X) # your API
def predict_proba(self, X):
return self.model.probabilities(X)
See the wrappers API reference for the full class documentation.