Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 27 additions & 29 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -40,35 +40,33 @@ The current paper version describes:
<!-- LEADERBOARD:START -->
| # | Estimator | Accuracy rank | Accuracy | Balanced accuracy | AUROC | F1 | Log loss &darr; | Sensitivity | Specificity |
|---|---|---|---|---|---|---|---|---|---|
| 1 | HC2 | **8.40** | **0.7887** | 0.7541 | **0.9000** | 0.7346 | **0.5440** | 0.7547 | **0.7910** |
| 2 | MRHydra | 9.21 | 0.7810 | **0.7579** | 0.8130 | **0.7368** | 7.8942 | **0.7715** | 0.7718 |
| 3 | RDST | 10.10 | 0.7707 | 0.7372 | 0.7963 | 0.7105 | 8.2660 | 0.7236 | 0.7833 |
| 4 | RIST | 10.75 | 0.7693 | 0.7422 | 0.8755 | 0.7221 | 0.6294 | 0.7504 | 0.7613 |
| 5 | DrCIF | 10.87 | 0.7721 | 0.7454 | 0.8821 | 0.7248 | 0.6558 | 0.7490 | 0.7669 |
| 6 | CIF | 11.07 | 0.7756 | 0.7497 | 0.8920 | 0.7288 | 0.6497 | 0.7536 | 0.7714 |
| 7 | FreshPRINCE | 11.08 | 0.7717 | 0.7516 | 0.8752 | 0.7293 | 0.6075 | 0.7515 | 0.7731 |
| 8 | Arsenal | 11.42 | 0.7654 | 0.7340 | 0.8471 | 0.7092 | 3.9265 | 0.7337 | 0.7696 |
| 9 | QUANT | 11.65 | 0.7693 | 0.7486 | 0.8839 | 0.7262 | 0.6238 | 0.7616 | 0.7539 |
| 10 | LITETime-MV | 11.97 | 0.7476 | 0.7312 | 0.8518 | 0.6875 | 1.3300 | 0.7200 | 0.7600 |
| 11 | ROCKET | 12.01 | 0.7661 | 0.7345 | 0.7955 | 0.7080 | 8.4299 | 0.7282 | 0.7724 |
| 12 | STSF | 12.60 | 0.7698 | 0.7503 | 0.8813 | 0.7155 | 0.6493 | 0.7439 | 0.7790 |
| 13 | H-InceptionTime | 12.71 | 0.7375 | 0.7205 | 0.8506 | 0.6897 | 1.3334 | 0.7303 | 0.7333 |
| 14 | LiteTIME | 13.22 | 0.7308 | 0.7122 | 0.8402 | 0.6746 | 1.4921 | 0.7199 | 0.7291 |
| 15 | DisjointCNN | 13.63 | 0.7286 | 0.7061 | 0.8354 | 0.6688 | 1.9705 | 0.6889 | 0.7368 |
| 16 | ConvTran | 14.37 | 0.7430 | 0.7139 | 0.8606 | 0.6882 | 0.8300 | 0.7289 | 0.7295 |
| 17 | Catch22 | 14.50 | 0.7442 | 0.7203 | 0.8703 | 0.6996 | 0.7238 | 0.7337 | 0.7326 |
| 18 | PatchMTSC | 14.51 | 0.7395 | 0.6934 | 0.8288 | 0.6660 | 0.7748 | 0.6985 | 0.7300 |
| 19 | STC | 15.19 | 0.7516 | 0.7188 | 0.8748 | 0.7004 | 0.6447 | 0.7264 | 0.7496 |
| 20 | TSF | 15.36 | 0.7484 | 0.7257 | 0.8747 | 0.6952 | 0.7335 | 0.7179 | 0.7565 |
| 21 | TS2Vec | 15.87 | 0.7212 | 0.6849 | 0.8082 | 0.6588 | 0.7326 | 0.6980 | 0.7100 |
| 22 | TDE | 15.93 | 0.7230 | 0.6823 | 0.8383 | 0.6441 | 0.8859 | 0.6786 | 0.7301 |
| 23 | Summary | 18.54 | 0.6814 | 0.6586 | 0.8263 | 0.6294 | 0.9251 | 0.6661 | 0.6787 |
| 24 | TimesNet | 18.86 | 0.6971 | 0.6688 | 0.8280 | 0.6390 | 1.1785 | 0.6850 | 0.6838 |
| 25 | TimesURL | 18.95 | 0.6916 | 0.6563 | 0.7931 | 0.6084 | 1.0193 | 0.6379 | 0.6914 |
| 26 | 1NN-DTW | 20.47 | 0.6672 | 0.6457 | 0.7214 | 0.6193 | 11.9949 | 0.6584 | 0.6584 |
| 27 | Dummy | 24.77 | 0.3538 | 0.2991 | 0.5000 | 0.1537 | 1.4284 | 0.2911 | 0.3695 |

Average over the 51 Multiverse-core datasets with results for every estimator on every metric, ordered by average accuracy rank. Best in each column in bold.
| 1 | HC2 | **7.64** | **0.7917** | **0.7557** | **0.8935** | **0.7302** | **0.5350** | 0.7469 | **0.7998** |
| 2 | MRHydra | 9.08 | 0.7794 | 0.7520 | 0.8040 | 0.7266 | 7.9526 | **0.7577** | 0.7768 |
| 3 | RDST | 9.46 | 0.7729 | 0.7386 | 0.7928 | 0.7075 | 8.1867 | 0.7172 | 0.7902 |
| 4 | RIST | 9.85 | 0.7744 | 0.7451 | 0.8679 | 0.7174 | 0.6150 | 0.7416 | 0.7743 |
| 5 | CIF | 10.23 | 0.7770 | 0.7487 | 0.8842 | 0.7246 | 0.6442 | 0.7459 | 0.7780 |
| 6 | DrCIF | 10.25 | 0.7731 | 0.7433 | 0.8745 | 0.7189 | 0.6430 | 0.7400 | 0.7736 |
| 7 | QUANT | 10.70 | 0.7668 | 0.7404 | 0.8694 | 0.7166 | 0.7285 | 0.7491 | 0.7570 |
| 8 | LITETime-MV | 10.95 | 0.7511 | 0.7320 | 0.8503 | 0.6878 | 1.3004 | 0.7167 | 0.7660 |
| 9 | Arsenal | 11.04 | 0.7663 | 0.7335 | 0.8419 | 0.7061 | 3.6444 | 0.7266 | 0.7752 |
| 10 | ROCKET | 11.09 | 0.7690 | 0.7362 | 0.7925 | 0.7065 | 8.3274 | 0.7228 | 0.7798 |
| 11 | STSF | 11.29 | 0.7727 | 0.7508 | 0.8723 | 0.7164 | 0.6685 | 0.7412 | 0.7845 |
| 12 | H-InceptionTime | 11.61 | 0.7421 | 0.7208 | 0.8447 | 0.6853 | 1.3448 | 0.7227 | 0.7436 |
| 13 | ConvTran | 12.88 | 0.7490 | 0.7177 | 0.8529 | 0.6862 | 0.8826 | 0.7234 | 0.7419 |
| 14 | DisjointCNN | 13.08 | 0.7296 | 0.7057 | 0.8246 | 0.6641 | 2.1100 | 0.6815 | 0.7431 |
| 15 | PatchMTSC | 13.37 | 0.7454 | 0.6990 | 0.8250 | 0.6671 | 0.7818 | 0.6981 | 0.7397 |
| 16 | Catch22 | 13.47 | 0.7463 | 0.7177 | 0.8605 | 0.6929 | 0.7068 | 0.7229 | 0.7420 |
| 17 | TSF | 14.01 | 0.7426 | 0.7175 | 0.8571 | 0.6896 | 0.9987 | 0.7095 | 0.7521 |
| 18 | STC | 14.03 | 0.7507 | 0.7137 | 0.8624 | 0.6803 | 0.6468 | 0.7036 | 0.7611 |
| 19 | TDE | 14.99 | 0.7251 | 0.6834 | 0.8339 | 0.6446 | 0.8524 | 0.6759 | 0.7349 |
| 20 | TS2Vec | 15.21 | 0.7201 | 0.6809 | 0.7994 | 0.6527 | 0.7835 | 0.6879 | 0.7144 |
| 21 | XCM | 16.58 | 0.6706 | 0.6370 | 0.7893 | 0.5771 | 2.1798 | 0.6167 | 0.6875 |
| 22 | Summary | 16.92 | 0.6845 | 0.6570 | 0.8113 | 0.6299 | 0.9645 | 0.6621 | 0.6852 |
| 23 | TimesNet | 17.24 | 0.7020 | 0.6709 | 0.8218 | 0.6396 | 1.2889 | 0.6810 | 0.6934 |
| 24 | TimesURL | 17.36 | 0.6950 | 0.6562 | 0.7831 | 0.6052 | 0.9868 | 0.6312 | 0.7016 |
| 25 | Dummy | 22.69 | 0.3709 | 0.3067 | 0.5000 | 0.1626 | 1.3928 | 0.2987 | 0.3880 |

Average over the 56 Multiverse-core datasets with results for every estimator on every metric, ordered by average accuracy rank. Best in each column in bold.
<!-- LEADERBOARD:END -->

Rebuilt with `python -m multiverse.experiments.tables`, which also writes a sortable
Expand Down
18 changes: 15 additions & 3 deletions docs/classifiers.md
Original file line number Diff line number Diff line change
Expand Up @@ -114,14 +114,26 @@ Mathematics, 9(23), 2021.
This is the only Keras port here, following the authors, so it needs `tensorflow`
rather than `torch`. Both are in the `deep-learning` extra.

The XCM results in this repository follow the authors' tuning protocol. Section 4.3 sets
**The XCM results published here are the fixed-parameter run**: a single fit at window
0.8 with batch 32, not the per-dataset search. The search was run over the full core 66
and did not pay for itself. Across the 65 shared datasets it was 0.017 mean accuracy
worse than the single fit, 31 wins to 31 with 3 ties, Wilcoxon p = 0.63. Worse, on the
14 datasets where the search happened to select 0.8 — the same window as the fixed run,
so the only difference is initialisation — the two runs still differed by 0.11 mean
absolute accuracy, and by 0.48 on HouseholdPowerConsumption2_disc and 0.44 on Libras.
At one resample XCM's run-to-run variance is larger than the effect the search is
tuning for, which makes the cross-validated selection largely a choice over noise: its
five per-window scores on Locust2022 run 0.674, 0.888, 0.253, 0.590, 0.707. The tuned
results are kept out of the tables rather than deleted.

The search itself follows the authors' protocol. Section 4.3 sets
`window_size` and `batch_size` per dataset "by grid search based on the best average
accuracy following a stratified 5-fold cross-validation on the training set", over
windows {0.2, 0.4, 0.6, 0.8, 1.0} and batches {1, 8, 32}. Selection never touches the
test data.

The reported run searches the window on that grid and holds batch size at 32. That is
the one departure, and it is a cost decision rather than a modelling one: batch 1 takes
That search holds batch size at 32 rather than searching it. That is a cost decision
rather than a modelling one: batch 1 takes
roughly 32 times the gradient steps, which would turn a day of GPU time into about 900
hours, for a value the published table selects on 4 of 30 datasets.

Expand Down
56 changes: 27 additions & 29 deletions docs/leaderboard.md
Original file line number Diff line number Diff line change
Expand Up @@ -109,35 +109,33 @@ inferred from the ranking.
<!-- UEA_LEADERBOARD:START -->
| # | Estimator | Accuracy rank | Accuracy | Balanced accuracy | AUROC | F1 | Log loss &darr; | Sensitivity | Specificity |
|---|---|---|---|---|---|---|---|---|---|
| 1 | HC2 | **7.37** | **0.7665** | **0.7452** | 0.8823 | 0.7411 | **0.6692** | 0.7470 | **0.7703** |
| 2 | RDST | 8.80 | 0.7459 | 0.7294 | 0.8179 | 0.7243 | 9.1587 | 0.7263 | 0.7560 |
| 3 | MRHydra | 9.20 | 0.7523 | 0.7388 | 0.8236 | **0.7432** | 8.9285 | 0.7599 | 0.7360 |
| 4 | Arsenal | 10.26 | 0.7321 | 0.7134 | 0.8439 | 0.7116 | 5.5624 | 0.7119 | 0.7417 |
| 5 | ROCKET | 10.28 | 0.7317 | 0.7146 | 0.8089 | 0.7130 | 9.6705 | 0.7133 | 0.7393 |
| 6 | RIST | 10.57 | 0.7433 | 0.7278 | 0.8755 | 0.7325 | 0.7983 | 0.7454 | 0.7322 |
| 7 | H-InceptionTime | 10.67 | 0.7223 | 0.7230 | 0.8653 | 0.6967 | 1.5030 | 0.7053 | 0.7345 |
| 8 | CIF | 11.33 | 0.7525 | 0.7378 | 0.8825 | 0.7400 | 0.8488 | **0.7604** | 0.7349 |
| 9 | FreshPRINCE | 11.74 | 0.7422 | 0.7281 | 0.8796 | 0.7239 | 0.7764 | 0.7313 | 0.7457 |
| 10 | DrCIF | 11.83 | 0.7386 | 0.7252 | 0.8734 | 0.7246 | 0.8458 | 0.7384 | 0.7303 |
| 11 | LITETime-MV | 12.00 | 0.7073 | 0.7064 | 0.8568 | 0.6779 | 1.4779 | 0.6905 | 0.7218 |
| 12 | LiteTIME | 12.50 | 0.7087 | 0.7019 | 0.8576 | 0.6751 | 1.6854 | 0.6985 | 0.7204 |
| 13 | DisjointCNN | 13.20 | 0.7011 | 0.7030 | 0.8510 | 0.6704 | 1.7943 | 0.6938 | 0.7070 |
| 14 | QUANT | 13.78 | 0.7285 | 0.7171 | **0.8888** | 0.7195 | 0.8041 | 0.7421 | 0.7074 |
| 15 | STSF | 14.24 | 0.7345 | 0.7223 | 0.8774 | 0.6934 | 0.8338 | 0.7007 | 0.7600 |
| 16 | TS2Vec | 14.78 | 0.7070 | 0.6913 | 0.8470 | 0.6917 | 0.8902 | 0.7150 | 0.6877 |
| 17 | TDE | 15.15 | 0.7079 | 0.6862 | 0.8484 | 0.6775 | 1.1475 | 0.6897 | 0.7095 |
| 18 | PatchMTSC | 15.37 | 0.7110 | 0.6986 | 0.8601 | 0.6899 | 0.7670 | 0.7192 | 0.6928 |
| 19 | ConvTran | 16.11 | 0.6931 | 0.6801 | 0.8552 | 0.6793 | 0.8155 | 0.7049 | 0.6736 |
| 20 | STC | 16.15 | 0.7265 | 0.7036 | 0.8803 | 0.7035 | 0.8124 | 0.7186 | 0.7184 |
| 21 | TSF | 16.17 | 0.7214 | 0.7076 | 0.8671 | 0.6917 | 0.9127 | 0.6977 | 0.7365 |
| 22 | Catch22 | 16.30 | 0.7006 | 0.6854 | 0.8557 | 0.6897 | 0.9814 | 0.7096 | 0.6802 |
| 23 | 1NN-DTW | 17.61 | 0.6848 | 0.6759 | 0.7785 | 0.6702 | 11.3600 | 0.6720 | 0.6879 |
| 24 | TimesURL | 18.24 | 0.6809 | 0.6658 | 0.8290 | 0.6539 | 1.3557 | 0.6698 | 0.6738 |
| 25 | Summary | 19.48 | 0.6477 | 0.6355 | 0.8295 | 0.6206 | 1.3000 | 0.6291 | 0.6589 |
| 26 | TimesNet | 20.09 | 0.6584 | 0.6504 | 0.8332 | 0.6386 | 1.1641 | 0.6628 | 0.6504 |
| 27 | Dummy | 24.78 | 0.2168 | 0.1980 | 0.5000 | 0.0800 | 1.9123 | 0.1853 | 0.2288 |

Average over the 23 UEA datasets with results for every estimator on every metric, ordered by average accuracy rank. Best in each column in bold.
| 1 | HC2 | **6.65** | **0.7617** | **0.7412** | 0.8752 | **0.7372** | **0.6681** | 0.7429 | **0.7655** |
| 2 | RDST | 8.27 | 0.7407 | 0.7250 | 0.8098 | 0.7197 | 9.3448 | 0.7212 | 0.7512 |
| 3 | MRHydra | 8.77 | 0.7462 | 0.7332 | 0.8145 | 0.7371 | 9.1480 | 0.7526 | 0.7315 |
| 4 | Arsenal | 9.38 | 0.7282 | 0.7103 | 0.8375 | 0.7084 | 5.3575 | 0.7084 | 0.7380 |
| 5 | H-InceptionTime | 9.56 | 0.7208 | 0.7214 | 0.8605 | 0.6962 | 1.5037 | 0.7046 | 0.7323 |
| 6 | ROCKET | 9.65 | 0.7265 | 0.7101 | 0.8004 | 0.7084 | 9.8595 | 0.7084 | 0.7341 |
| 7 | RIST | 10.04 | 0.7374 | 0.7226 | 0.8660 | 0.7261 | 0.7933 | 0.7370 | 0.7292 |
| 8 | CIF | 10.42 | 0.7475 | 0.7334 | 0.8742 | 0.7356 | 0.8410 | **0.7550** | 0.7307 |
| 9 | LITETime-MV | 10.96 | 0.7048 | 0.7040 | 0.8505 | 0.6769 | 1.4815 | 0.6894 | 0.7181 |
| 10 | DrCIF | 11.10 | 0.7328 | 0.7199 | 0.8639 | 0.7193 | 0.8386 | 0.7324 | 0.7250 |
| 11 | QUANT | 12.54 | 0.7245 | 0.7136 | **0.8803** | 0.7161 | 0.7978 | 0.7383 | 0.7035 |
| 12 | DisjointCNN | 12.62 | 0.6943 | 0.6961 | 0.8387 | 0.6627 | 1.8839 | 0.6830 | 0.7044 |
| 13 | STSF | 12.71 | 0.7309 | 0.7191 | 0.8703 | 0.6914 | 0.8255 | 0.6982 | 0.7556 |
| 14 | PatchMTSC | 13.71 | 0.7096 | 0.6977 | 0.8549 | 0.6906 | 0.7619 | 0.7220 | 0.6874 |
| 15 | TS2Vec | 13.96 | 0.6990 | 0.6839 | 0.8334 | 0.6853 | 0.8820 | 0.7088 | 0.6783 |
| 16 | TDE | 14.19 | 0.7026 | 0.6818 | 0.8386 | 0.6745 | 1.1277 | 0.6877 | 0.7016 |
| 17 | ConvTran | 14.23 | 0.6915 | 0.6791 | 0.8493 | 0.6784 | 0.8090 | 0.7032 | 0.6725 |
| 18 | TSF | 14.50 | 0.7183 | 0.7052 | 0.8600 | 0.6901 | 0.9018 | 0.6962 | 0.7324 |
| 19 | STC | 14.62 | 0.7224 | 0.7004 | 0.8722 | 0.7000 | 0.8052 | 0.7138 | 0.7157 |
| 20 | Catch22 | 14.98 | 0.6945 | 0.6800 | 0.8443 | 0.6836 | 0.9690 | 0.7021 | 0.6762 |
| 21 | XCM | 15.98 | 0.6356 | 0.6204 | 0.8184 | 0.5899 | 1.3507 | 0.6172 | 0.6410 |
| 22 | TimesURL | 17.04 | 0.6745 | 0.6600 | 0.8170 | 0.6473 | 1.3326 | 0.6612 | 0.6703 |
| 23 | TimesNet | 17.92 | 0.6582 | 0.6505 | 0.8275 | 0.6400 | 1.1485 | 0.6648 | 0.6481 |
| 24 | Summary | 18.06 | 0.6431 | 0.6314 | 0.8181 | 0.6179 | 1.2745 | 0.6269 | 0.6521 |
| 25 | Dummy | 23.15 | 0.2286 | 0.2106 | 0.5000 | 0.1044 | 1.8615 | 0.2193 | 0.2193 |

Average over the 24 UEA datasets with results for every estimator on every metric, ordered by average accuracy rank. Best in each column in bold.
<!-- UEA_LEADERBOARD:END -->

Sortable version with per-metric ranks:
Expand Down
81 changes: 63 additions & 18 deletions multiverse/classification/_rankscl.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,20 @@ def _ranking_loss(embeddings, labels, distance):

Returns None when the batch has no anchor with both a positive and a
negative, which the authors' version would raise on.

Evaluated as one dense expression rather than the authors' loop over anchors
and positives. The loop is what the paper describes, but at the archive
settings a batch holds ``batch_size * (2 * aug_positives + 1)`` embeddings,
44 by default, and every case is repeated, so each anchor has around ten
positives. That is roughly 440 Python iterations per optimiser step, each
launching several small CUDA kernels, and launch latency then decides the
runtime rather than the arithmetic: measured at about 3.5 seconds a step for
a small encoder on an H200, which timed LSST out of a 60 hour job at epoch
91 of 100. The value is unchanged; only the order of the reduction differs.

Memory is cubic in the batch, ``n * n * n`` floats for the pairwise
differences, which is 85k elements at the defaults. A much larger
``batch_size`` or ``aug_positives`` would need this chunked over anchors.
"""
if distance == "Cosine":
matrix = -torch.cosine_similarity(
Expand All @@ -185,20 +199,23 @@ def _ranking_loss(embeddings, labels, distance):
matrix = torch.cdist(embeddings, embeddings, p=2)

same = labels.reshape(1, -1) == labels.reshape(-1, 1)
violations = []
for anchor in range(matrix.shape[0]):
negatives = matrix[anchor][~same[anchor]]
if negatives.numel() == 0:
continue
positives = same[anchor].nonzero().flatten()
for positive in positives[positives != anchor]:
gap = matrix[anchor, positive]
closer = negatives[negatives <= gap]
violations.append(torch.sigmoid(gap - closer).sum())

if not violations:
negative = ~same
# a positive is a same-class case other than the anchor, and an anchor with
# no negative is skipped entirely, as in the authors' loop
pairs = same & ~torch.eye(
matrix.shape[0], dtype=torch.bool, device=matrix.device
)
pairs = pairs & negative.any(dim=1, keepdim=True)
if not pairs.any():
return None
return torch.atan(torch.stack(violations)).mean()

# difference[anchor, positive, other] = d(anchor, positive) - d(anchor, other)
difference = matrix.unsqueeze(2) - matrix.unsqueeze(1)
# the wrongly ranked ones: a negative at least as close as the positive is
wrong = negative.unsqueeze(1) & (difference >= 0)
violations = (torch.sigmoid(difference) * wrong).sum(dim=2)

return torch.atan(violations)[pairs].mean()


class RankSCLClassifier(BaseClassifier):
Expand Down Expand Up @@ -359,15 +376,23 @@ def _subsample(self, features, y):
return features, y

def _build_probe(self, n_cases: int, seed: int):
"""Return the probe, following ``utils/_eval_protocols.py``."""
"""Return the probe, following ``utils/_eval_protocols.py``.

The SVM is built with ``probability=False``, which is what the authors'
grid sets. Platt scaling is not free: libsvm fits it by an internal
five-fold cross-validation inside every ``fit``, so an estimator carrying
``probability=True`` into a ten-value grid over five folds costs about
300 SVC trainings rather than 50. ``_fit_probe`` turns it back on for a
single refit at the selected C, which is what ``predict_proba`` needs.
"""
if self.probe == "logistic":
return make_pipeline(
StandardScaler(),
OneVsRestClassifier(
LogisticRegression(max_iter=1000000, random_state=seed)
),
)
svm = SVC(C=np.inf, gamma="scale", probability=True, random_state=seed)
svm = SVC(C=np.inf, gamma="scale", probability=False, random_state=seed)
if n_cases // self.n_classes_ < 5 or n_cases < 50:
return svm
return GridSearchCV(
Expand All @@ -378,6 +403,28 @@ def _build_probe(self, n_cases: int, seed: int):
n_jobs=1,
)

def _fit_probe(self, features, y, seed):
"""Select the probe's parameters, then refit it with probabilities on.

The selection is the authors' own: their grid sets
``probability=False``, and scoring uses ``predict``, which reads the
decision function either way, so the chosen C is unchanged. Only the
final estimator needs Platt scaling, because aeon classifiers must
implement ``predict_proba``.
"""
probe = self._build_probe(features.shape[0], seed)
if self.probe == "logistic":
return probe.fit(features, y)
if isinstance(probe, GridSearchCV):
probe.fit(features, y)
parameters = probe.best_params_
else:
# the degenerate-case bypass: too few cases per class to select on
parameters = {"C": probe.C, "kernel": "rbf", "gamma": "scale"}
return SVC(probability=True, random_state=seed, **parameters).fit(
features, y
)

def _fit(self, X: np.ndarray, y):
self._validate_parameters()

Expand Down Expand Up @@ -452,9 +499,7 @@ def _fit(self, X: np.ndarray, y):
representations = self._encode(X)
fit_features, fit_y = self._subsample(representations, encoded_y)
self.probe_cases_ = int(fit_features.shape[0])
self.probe_ = self._build_probe(self.probe_cases_, seed).fit(
fit_features, fit_y
)
self.probe_ = self._fit_probe(fit_features, fit_y, seed)
return self

def _check_shape(self, X: np.ndarray) -> None:
Expand Down
Loading
Loading