ASTRAL User Guide & Practical API Manual
A hands-on manual for data scientists and ML engineers using ASTRAL for tabular classification, regression, conformal prediction, and causal-invariant feature interpretation.
01 Installation
ASTRAL is a pure-NumPy tabular learning library with zero mandatory external C++ or deep learning compiler dependencies.
# Install official package from PyPI
pip install astral-model
# Or install with optional extras (Torch acceleration, Scikit-Learn benchmarking)
pip install "astral-model[all]"
# Editable installation from local source repository
pip install -e .
02 5-Minute Quickstart
ASTRAL follows the familiar Scikit-Learn estimator pattern (fit, predict, predict_proba, score):
from astral import AstralModel, AstralScaler, astral_train_test_split
from sklearn.datasets import load_breast_cancer
# 1. Load data
X, y = load_breast_cancer(return_X_y=True)
# 2. Train/Test split with stratification
X_train, X_test, y_train, y_test = astral_train_test_split(
X, y, test_size=0.2, random_state=42, stratify=y
)
# 3. Fit scaler & model
scaler = AstralScaler()
X_train_s = scaler.fit_transform(X_train)
X_test_s = scaler.transform(X_test)
model = AstralModel(ridge="auto", random_state=42)
model.fit(X_train_s, y_train)
# 4. Predict & evaluate
y_pred = model.predict(X_test_s)
metrics = model.score(X_test_s, y_test)
print(f"Test Accuracy: {metrics['accuracy']:.4f}, Macro F1: {metrics['f1_macro']:.4f}")
03 Constructor Parameters Reference
Detailed breakdown of all parameters accepted by AstralModel(...):
| Parameter | Type / Default | Description & Guidance |
|---|---|---|
| task | str = 'auto' | 'auto', 'classification', or 'regression'. In auto mode, inferred from target discrete/continuous characteristics. |
| n_quantiles | int = 5 | Number of empirical quantile anchor points per feature for hat, RBF, and threshold basis placements. Default 5 is optimal for most tabular datasets. |
| max_bases_per_feature | int = 6 | Maximum univariate non-linear basis functionals generated per raw marginal feature during Phase 2. |
| max_interactions | int = 35 | Maximum pairwise resonance interactions (multiplicative, ratio, min, diff) retained in the design dictionary. |
| n_environments | int = 10 | Number of synthetic bootstrap perturbation environments generated to evaluate Invariant Risk Minimization (IRM) stability. |
| stability_threshold | float = 0.5 | Minimum IRM stability score s_j required for a feature to be considered structurally stable and receive reduced shrinkage. |
| alpha | float = 0.1 | Significance level for distribution-free conformal uncertainty intervals and sets (default 0.1 provides 90% coverage guarantee). |
| ridge | float | str = 'auto' | L2 regularization parameter λ. When 'auto', automatically tuned via fast stratified K-fold cross-validation. |
| ridge_grid | list | None = None | Custom candidate regularizer values to evaluate during cross-validation grid search. Defaults to logarithmic mesh [1e-4, 1e-3, ..., 1e3]. |
| class_weight | str | dict | None = None | 'balanced' or custom class-to-weight dictionary to counter severe target class imbalance. |
| impute_missing | bool = True | When True, automatically replaces NaNs and infinite values with median statistics computed during training. |
| max_samples_discovery | int = 4000 | Maximum sub-sample size used for basis synthesis and IRM ranking on large datasets, guaranteeing sub-3-second discovery. |
| random_state | int = 42 | Random seed for reproducible environment bootstrap resampling and CV splits. |
| verbose | bool = False | If True, logs progress of basis discovery, IRM stability filtering, and regularizer selection. |
04 Methods & API Reference
| Method | Signature | Returns & Functionality |
|---|---|---|
fit(X, y) |
(X, y) |
Fits the 6-phase ASTRAL model pipeline. Returns self. |
predict(X) |
(X) → np.ndarray |
Predicts class labels for classification or scalar continuous targets for regression. |
predict_proba(X) |
(X) → np.ndarray |
Returns calibrated class probability distributions of shape (N, C). |
predict_set(X, alpha) |
(X, alpha=None, allow_abstention=True) → list[list] |
Returns conformal prediction sets for each sample with guaranteed 1 - α coverage. |
predict_with_uncertainty(X, alpha) |
(X, alpha=None) → (y_pred, y_lower, y_upper) |
Returns point predictions and exact conformal prediction interval bounds. |
score(X, y) |
(X, y) → dict |
Computes accuracy, macro F1, weighted F1, and ROC-AUC (classification) or R², RMSE, MAE (regression). |
get_feature_importance() |
() → list[dict] |
Returns ranked basis importance scores alongside IRM stability flags. |
get_stable_features() |
() → list[str] |
Returns names of all features verified invariant across bootstrap perturbation environments. |
explain(x, top_k) |
(x, top_k=5) → dict |
Generates local linear basis decomposition explaining an individual instance prediction. |
summary() |
() → str |
Generates formatted textual summary of active bases, IRM stability, and learned weights. |
data_quality_report() |
() → dict |
Inspects NaN rates, infinite values, constant columns, and imputation summary. |
05 Classification with Conformal Prediction Sets
from astral import AstralModel, astral_train_test_split
import numpy as np
# Train model
model = AstralModel(task="classification", ridge="auto", random_state=42)
model.fit(X_train, y_train)
# Conformal prediction sets with 90% guaranteed coverage (alpha=0.1)
pred_sets = model.predict_set(X_test, alpha=0.1)
# Inspect sample predictions
for i in range(5):
probs = model.predict_proba(X_test[i:i+1])[0]
print(f"Sample {i}: Pred={model.predict(X_test[i:i+1])[0]}, Set={pred_sets[i]}, Probs={np.round(probs, 3)}")
06 Regression & Uncertainty Intervals
model = AstralModel(task="regression", ridge="auto", random_state=42)
model.fit(X_train, y_train)
# Predict with calibrated conformal interval
y_pred, y_lower, y_upper = model.predict_with_uncertainty(X_test, alpha=0.05) # 95% coverage
# Calculate empirical coverage
coverage = np.mean((y_test >= y_lower) & (y_test <= y_upper))
print(f"Empirical Coverage: {coverage * 100:.2f}% (Target: 95%)")
print(f"Average Interval Width: {np.mean(y_upper - y_lower):.2f}")
07 Model Explainability & Invariance Inspection
# 1. Inspect global basis importances and causal invariance
for item in model.get_feature_importance()[:8]:
status = "[INVARIANT]" if item["is_stable"] else "[UNSTABLE]"
print(f"{status:<12} {item['name']:<30} Importance: {item['importance']:.4f}")
# 2. Local instance explanation
explanation = model.explain(X_test[0], top_k=3)
print("\nLocal Instance Prediction Breakdown:")
for contribution in explanation["contributions"]:
print(f" {contribution['basis_name']}: contribution = {contribution['value']:+.4f}")
08 Empirical Benchmark Results
Empirical performance across three benchmark datasets compared to 10 mainstream algorithms:
| Dataset | Task & Samples | ASTRAL | Gradient Boosting | Random Forest | Logistic / Ridge |
|---|---|---|---|---|---|
| data-1 (Ames Housing) | Regression (1,460 rows, 79 cols) | R² = 0.8596 (RMSE: 32,819) | R² = 0.9031 | R² = 0.8830 | R² = 0.8070 |
| data-2 (EV Adoption) | Classification (15,000 samples) | AUC = 0.9394 (Acc: 89.4%) | AUC = 0.9376 | AUC = 0.9336 | AUC = 0.8725 (raw) |
| data-3 (Heart Disease) | Classification (303 samples) | AUC = 0.9044 (Acc: 86.9%) | AUC = 0.8821 | AUC = 0.8936 | AUC = 0.8853 |
09 Frequently Asked Questions (FAQ)
Q: Do I need to normalize or scale features before training?
While AstralScaler is provided for convenience, ASTRAL is intrinsically robust to unscaled heterogeneous data because its non-linear basis transforms are positioned directly at empirical training quantiles.
Q: How does ASTRAL handle missing values (NaNs)?
By default (impute_missing=True), ASTRAL automatically detects missing values, computes training column medians, and fills missing inputs at both fit and predict time without data leakage.
Q: Can ASTRAL be pickled or saved with Joblib?
Yes. ASTRAL uses pure standard Python data structures and NumPy arrays, making it 100% serializable via standard pickle or joblib.dump().