ModelEvaluator.run
- ModelEvaluator.run(X, labels, n_cv=5, n_rounds=1, metrics=None, ci=0.95, random_state=None)[source]
Evaluate every model by repeated stratified cross-validation with bootstrap CIs.
Runs
n_roundsrepeats of stratifiedn_cv-fold cross-validation (each repeat reshuffled with a distinct seed derived fromrandom_state), scoring every model on the same folds. The per-fold scores are aggregated per (model, metric) into a mean, a population std, a percentile bootstrap confidence interval of the mean, and the fold count. The raw per-fold scores are stored ondf_scores_foreval().Added in version 1.1.0.
- Parameters:
X (array-like, shape (n_samples, n_features)) – Feature matrix.
labels (array-like, shape (n_samples,)) – Binary class labels for the samples in
X.n_cv (int, default=5) – Number of stratified cross-validation folds per round (must not exceed the smallest class count).
n_rounds (int, default=1) – Number of cross-validation repeats (multi-seed aggregation). The total number of fold scores per (model, metric) is
n_cv * n_rounds.metrics (list of str, optional) – Performance metrics to compute. Defaults to
list_metricsfrom the constructor.ci (float, optional) – Central confidence level in
(0, 1)for the percentile bootstrap CI of the mean. IfNone, theci_low/ci_highcolumns areNaN. Default is0.95.random_state (int, optional) – Per-call seed overriding the constructor’s
random_statefor this evaluation.
- Returns:
df_eval – Long-format evaluation table with columns
model,metric,score(mean over folds),score_std(population std over folds),ci_low/ci_high(bootstrap CI of the mean,NaNwhenciisNone), andn_scores(fold count).- Return type:
pd.DataFrame, shape (n_models * n_metrics, 7)
Examples
runscores every model by repeated stratified cross-validation and returns a per (model, metric) table with mean, std, bootstrap CI, and fold count. First, theDOM_GSECdataset and its feature matrix:import aaanalysis as aa aa.options["verbose"] = False # Disable verbosity # DOM_GSEC example dataset + a small feature set (see [Breimann25]_) df_seq = aa.load_dataset(name="DOM_GSEC") labels = df_seq["label"].to_list() df_feat = aa.load_features(name="DOM_GSEC").head(20) # Build the CPP feature matrix X sf = aa.SequenceFeature() df_parts = sf.get_df_parts(df_seq=df_seq) X = sf.feature_matrix(features=df_feat["feature"], df_parts=df_parts)
Pass
Xandlabelswith the number of folds (n_cv), the number of cross-validation repeats (n_rounds, multi-seed aggregation), themetrics, the bootstrap confidence levelci, and a per-callrandom_state:me = aa.ModelEvaluator(models=["rf", "svm"], random_state=42, verbose=False) df_eval = me.run(X=X, labels=labels, n_cv=5, n_rounds=3, metrics=["balanced_accuracy", "mcc", "roc_auc"], ci=0.95, random_state=42) aa.display_df(df_eval, n_rows=10, show_shape=True)
DataFrame shape: (6, 7)
/Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn( /Users/stephanbreimann/Programming/1Packages/aaanalysis/.venv/lib/python3.13/site-packages/sklearn/svm/_base.py:239: FutureWarning: The probability parameter was deprecated in 1.9 and will be removed in version 1.11. Use CalibratedClassifierCV(SVC(), ensemble=False) instead of SVC(probability=True) warnings.warn(
model metric score score_std ci_low ci_high n_scores 1 rf balanced_accuracy 0.803632 0.076990 0.763862 0.838034 15 2 rf mcc 0.612956 0.151909 0.534722 0.680772 15 3 rf roc_auc 0.890976 0.071330 0.852509 0.923964 15 4 svm balanced_accuracy 0.848504 0.094490 0.799995 0.893803 15 5 svm mcc 0.703752 0.186625 0.606794 0.793224 15 6 svm roc_auc 0.908876 0.064904 0.876017 0.940368 15