ASTRAL User Guide & Reference v0.2.1
Switch to Developer Guide →

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().