pyensemblefs 0.3.13__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (67) hide show
  1. pyensemblefs/__init__.py +17 -0
  2. pyensemblefs/aggregators/__init__.py +99 -0
  3. pyensemblefs/aggregators/abcvote.py +270 -0
  4. pyensemblefs/aggregators/base.py +170 -0
  5. pyensemblefs/aggregators/rank.py +158 -0
  6. pyensemblefs/aggregators/score.py +219 -0
  7. pyensemblefs/aggregators/subset.py +310 -0
  8. pyensemblefs/api.py +212 -0
  9. pyensemblefs/cli.py +52 -0
  10. pyensemblefs/configs.py +33 -0
  11. pyensemblefs/datasets/__init__.py +13 -0
  12. pyensemblefs/datasets/breast_cancer.py +9 -0
  13. pyensemblefs/datasets/cancer.py +9 -0
  14. pyensemblefs/datasets/diabetes_sklearn.py +12 -0
  15. pyensemblefs/datasets/heart.py +13 -0
  16. pyensemblefs/datasets/pima.py +16 -0
  17. pyensemblefs/ensemble/__init__.py +0 -0
  18. pyensemblefs/ensemble/base.py +219 -0
  19. pyensemblefs/ensemble/bootstrapper.py +94 -0
  20. pyensemblefs/ensemble/featureselector.py +63 -0
  21. pyensemblefs/ensemble/metabootstrapper.py +256 -0
  22. pyensemblefs/estimators/base.py +116 -0
  23. pyensemblefs/estimators/evaluator.py +61 -0
  24. pyensemblefs/fsmethods/__init__.py +20 -0
  25. pyensemblefs/fsmethods/basefs.py +164 -0
  26. pyensemblefs/fsmethods/factory.py +146 -0
  27. pyensemblefs/fsmethods/rank.py +124 -0
  28. pyensemblefs/fsmethods/score.py +267 -0
  29. pyensemblefs/fsmethods/subset.py +308 -0
  30. pyensemblefs/fsmethods/variance.py +10 -0
  31. pyensemblefs/main.py +62 -0
  32. pyensemblefs/main_stab.py +173 -0
  33. pyensemblefs/pipeline.py +43 -0
  34. pyensemblefs/selectors/__init__.py +22 -0
  35. pyensemblefs/selectors/base.py +45 -0
  36. pyensemblefs/selectors/filters.py +111 -0
  37. pyensemblefs/selectors/model_based.py +118 -0
  38. pyensemblefs/stability/__init__.py +1 -0
  39. pyensemblefs/stability/base.py +26 -0
  40. pyensemblefs/stability/config.py +35 -0
  41. pyensemblefs/stability/evaluator.py +181 -0
  42. pyensemblefs/stability/expectations.py +62 -0
  43. pyensemblefs/stability/frequency.py +16 -0
  44. pyensemblefs/stability/helpers.py +120 -0
  45. pyensemblefs/stability/measures_adjusted_intersections.py +110 -0
  46. pyensemblefs/stability/measures_adjusted_other.py +122 -0
  47. pyensemblefs/stability/measures_unadjusted.py +175 -0
  48. pyensemblefs/stability/pairwise.py +21 -0
  49. pyensemblefs/stability/stability.py +283 -0
  50. pyensemblefs/stability/utils_io.py +48 -0
  51. pyensemblefs/tools/generate_feature_sets.py +131 -0
  52. pyensemblefs/tools/sim_matrix.py +86 -0
  53. pyensemblefs/utils/consts.py +13 -0
  54. pyensemblefs/utils/datasets.py +55 -0
  55. pyensemblefs/utils/loader.py +18 -0
  56. pyensemblefs/utils/plotter.py +24 -0
  57. pyensemblefs/viz/__init__.py +0 -0
  58. pyensemblefs/viz/comparison.py +134 -0
  59. pyensemblefs/viz/ranking.py +79 -0
  60. pyensemblefs/viz/stability.py +107 -0
  61. pyensemblefs/viz/visualizer.py +23 -0
  62. pyensemblefs-0.3.13.dist-info/METADATA +293 -0
  63. pyensemblefs-0.3.13.dist-info/RECORD +67 -0
  64. pyensemblefs-0.3.13.dist-info/WHEEL +5 -0
  65. pyensemblefs-0.3.13.dist-info/entry_points.txt +6 -0
  66. pyensemblefs-0.3.13.dist-info/licenses/LICENSE +21 -0
  67. pyensemblefs-0.3.13.dist-info/top_level.txt +1 -0
@@ -0,0 +1,17 @@
1
+ from __future__ import annotations
2
+ from importlib.metadata import version, PackageNotFoundError
3
+
4
+ from . import datasets
5
+ from .api import get_config, compute_scores, extract_features
6
+
7
+ __all__ = [
8
+ "aggregators","ensemble","estimators","fsmethods","selectors",
9
+ "stability","tools","utils","viz",
10
+ "datasets","get_config","compute_scores","extract_features"
11
+ ]
12
+
13
+ def _get_version() -> str:
14
+ try: return version("pyensemblefs")
15
+ except PackageNotFoundError: return "0.0.0"
16
+
17
+ __version__ = _get_version()
@@ -0,0 +1,99 @@
1
+ # -*- coding: utf-8 -*-
2
+ """
3
+ Unified public API for aggregators in pyensemblefs.
4
+
5
+ This module re-exports the most common aggregators from the score-, rank-, subset-
6
+ and ABC-voting modules, so users can simply do:
7
+
8
+ from pyensemblefs.aggregators import (
9
+ # Base types
10
+ ScoreAggregator, RankAggregator, BinaryAggregator,
11
+ # Score-based
12
+ MeanAggregator, SumAggregator, MedianAggregator, WeightedScoreAggregator,
13
+ BordaFromScoresAggregator, SelectionFrequencyAggregator,
14
+ # Rank-based
15
+ MeanRankAggregator, MedianRankAggregator, ConsensusRankAggregator,
16
+ BordaFromRanksAggregator, TrimmedMeanRankAggregator,
17
+ WinsorizedMeanRankAggregator, GeometricMeanRankAggregator, RankProductAggregator,
18
+ # Subset-based
19
+ ThresholdAggregator, MajorityVoteAggregator, TopKBinaryAggregator,
20
+ # ABC-voting (generic + factory + catalog)
21
+ ABCVotingRule, make_abcvoter, SAFE_RULES,
22
+ )
23
+ """
24
+
25
+ # ABC-voting (generic rule + factory + catalog)
26
+ from .abcvote import (
27
+ ABCVotingRule,
28
+ make_abcvoter,
29
+ SAFE_RULES,
30
+ )
31
+
32
+ # Base classes
33
+ from .base import BaseAggregator, ScoreAggregator, RankAggregator, BinaryAggregator
34
+
35
+ # Score-based aggregators
36
+ from .score import (
37
+ MeanAggregator,
38
+ SumAggregator,
39
+ MedianAggregator,
40
+ WeightedScoreAggregator,
41
+ BordaFromScoresAggregator,
42
+ SelectionFrequencyAggregator,
43
+ )
44
+
45
+ # Rank-based aggregators
46
+ from .rank import (
47
+ MeanRankAggregator,
48
+ MedianRankAggregator,
49
+ ConsensusRankAggregator,
50
+ BordaFromRanksAggregator,
51
+ TrimmedMeanRankAggregator,
52
+ WinsorizedMeanRankAggregator,
53
+ GeometricMeanRankAggregator,
54
+ RankProductAggregator,
55
+ )
56
+
57
+ # Subset-based aggregators
58
+ from .subset import (
59
+ ThresholdAggregator,
60
+ MajorityVoteAggregator,
61
+ TopKBinaryAggregator,
62
+ )
63
+
64
+ __all__ = [
65
+ # base
66
+ "BaseAggregator",
67
+ "ScoreAggregator",
68
+ "RankAggregator",
69
+ "BinaryAggregator",
70
+
71
+ # score-based
72
+ "MeanAggregator",
73
+ "SumAggregator",
74
+ "MedianAggregator",
75
+ "WeightedScoreAggregator",
76
+ "BordaFromScoresAggregator",
77
+ "SelectionFrequencyAggregator",
78
+
79
+ # rank-based
80
+ "MeanRankAggregator",
81
+ "MedianRankAggregator",
82
+ "ConsensusRankAggregator",
83
+ "BordaFromRanksAggregator",
84
+ "TrimmedMeanRankAggregator",
85
+ "WinsorizedMeanRankAggregator",
86
+ "GeometricMeanRankAggregator",
87
+ "RankProductAggregator",
88
+
89
+ # subset-based
90
+ "ThresholdAggregator",
91
+ "MajorityVoteAggregator",
92
+ "TopKBinaryAggregator",
93
+
94
+ # ABC-voting (generic only)
95
+ "ABCVotingRule",
96
+ "make_abcvoter",
97
+ "SAFE_RULES",
98
+ ]
99
+
@@ -0,0 +1,270 @@
1
+ from __future__ import annotations
2
+
3
+ import numpy as np
4
+ import pandas as pd
5
+ from typing import Callable, Dict, List, Set, Optional
6
+
7
+ from .base import ScoreAggregator
8
+ from abcvoting.preferences import Profile
9
+ from abcvoting import abcrules as _abcr
10
+ from abcvoting.abcrules import Rule
11
+
12
+
13
+ def _committee_ids_from_abcvoting_object(obj) -> Set[int]:
14
+ if isinstance(obj, (list, tuple)) and len(obj) > 0 and not isinstance(obj, set):
15
+ obj = obj[0]
16
+
17
+ try:
18
+ return set(int(x) for x in obj)
19
+ except Exception:
20
+ pass
21
+
22
+ try:
23
+ return set(int(getattr(x, "id")) for x in obj)
24
+ except Exception:
25
+ pass
26
+
27
+ if hasattr(obj, "aslist"):
28
+ try:
29
+ return set(int(x) for x in obj.aslist())
30
+ except Exception:
31
+ pass
32
+
33
+ if isinstance(obj, set):
34
+ try:
35
+ return set(int(x) for x in list(obj))
36
+ except Exception:
37
+ pass
38
+
39
+ raise TypeError(
40
+ f"Cannot normalize abcvoting committee object of type {type(obj)} to a set of ints."
41
+ )
42
+
43
+
44
+ def _pick_safe_algorithm(rule_id: str, prefer: Optional[List[str]] = None) -> Optional[str]:
45
+ """Return an algorithm for rule_id that does not depend on Gurobi."""
46
+ try:
47
+ algos = tuple(a for a in Rule(rule_id).algorithms if "gurobi" not in a.lower())
48
+ except Exception:
49
+ return None
50
+ if not algos:
51
+ return None
52
+ prefer = prefer or []
53
+ for a in prefer:
54
+ if a in algos:
55
+ return a
56
+ return algos[0]
57
+
58
+
59
+
60
+ class _ABCBase(ScoreAggregator):
61
+ """Base class for ABC voting aggregators using `abcvoting`."""
62
+
63
+ rule_name_: str = "ABC-Rule"
64
+
65
+ def __init__(self, top_k: int):
66
+ super().__init__(top_k=top_k)
67
+ if top_k is None or int(top_k) <= 0:
68
+ raise ValueError("top_k (committee size) must be a positive integer.")
69
+ self._committee_: Set[int] | None = None
70
+
71
+ def _build_profile(self, R: np.ndarray) -> Profile:
72
+ if R.ndim != 2:
73
+ raise ValueError("results must be a 2D array (n_bootstraps, n_features).")
74
+ if not np.array_equal(R, R.astype(int)):
75
+ raise ValueError(f"{self.rule_name_} expects binary 0/1 results.")
76
+ n_boot, n_feat = R.shape
77
+ approvals = [set(np.where(R[b] == 1)[0].tolist()) for b in range(n_boot)]
78
+ profile = Profile(n_feat)
79
+ profile.add_voters(approvals)
80
+ return profile
81
+
82
+ @staticmethod
83
+ def _scores_av(R: np.ndarray) -> np.ndarray:
84
+ return R.sum(axis=0).astype(float)
85
+
86
+ @staticmethod
87
+ def _scores_sav(R: np.ndarray) -> np.ndarray:
88
+ row_sums = R.sum(axis=1)
89
+ scores = np.zeros(R.shape[1], dtype=float)
90
+ for b in range(R.shape[0]):
91
+ m = int(row_sums[b])
92
+ if m > 0:
93
+ scores += (R[b] / float(m))
94
+ return scores
95
+
96
+ _scores_slav = _scores_sav # same logic
97
+
98
+ def aggregate(self, results: np.ndarray) -> np.ndarray:
99
+ raise NotImplementedError
100
+
101
+ def fit(self, results: np.ndarray):
102
+ agg_result = self.aggregate(results)
103
+ sorted_indices = np.argsort(-agg_result)
104
+ self.final_ranking_ = sorted_indices[: self.top_k]
105
+ n_feat = results.shape[1]
106
+ mask = np.zeros(n_feat, dtype=int)
107
+ if self._committee_ is None:
108
+ mask[self.final_ranking_] = 1
109
+ else:
110
+ for j in self._committee_:
111
+ mask[int(j)] = 1
112
+ self.selected_features_ = mask
113
+ self._agg_result = agg_result
114
+ return self
115
+
116
+
117
+
118
+ class ABCVotingRule(_ABCBase):
119
+ """Generic aggregator for any `abcvoting` rule (without Gurobi)."""
120
+
121
+ def __init__(
122
+ self,
123
+ top_k: int,
124
+ compute_func: Callable,
125
+ rule_id: Optional[str] = None,
126
+ score_mode: str = "av",
127
+ prefer_algorithms: Optional[List[str]] = None,
128
+ **extra_kwargs,
129
+ ):
130
+ super().__init__(top_k=top_k)
131
+ self.compute_func = compute_func
132
+ self.rule_id = rule_id
133
+ self.score_mode = score_mode.lower()
134
+ self.prefer_algorithms = prefer_algorithms or []
135
+ self.extra_kwargs = {"resolute": True, **extra_kwargs}
136
+ self.rule_name_ = getattr(compute_func, "__name__", rule_id or "ABC-Rule")
137
+
138
+ def aggregate(self, results: np.ndarray) -> np.ndarray:
139
+ R = np.asarray(results)
140
+ profile = self._build_profile(R)
141
+ if "algorithm" not in self.extra_kwargs and self.rule_id is not None:
142
+ safe_alg = _pick_safe_algorithm(self.rule_id, self.prefer_algorithms)
143
+ if safe_alg is not None:
144
+ self.extra_kwargs["algorithm"] = safe_alg
145
+ committees = self.compute_func(profile, committeesize=int(self.top_k), **self.extra_kwargs)
146
+ self._committee_ = _committee_ids_from_abcvoting_object(committees)
147
+ if self.score_mode == "sav":
148
+ return self._scores_sav(R)
149
+ if self.score_mode == "slav":
150
+ return self._scores_slav(R)
151
+ return self._scores_av(R)
152
+
153
+
154
+
155
+
156
+
157
+ SAFE_RULES: Dict[str, Dict] = {
158
+ "av": {"rule_id": "av", "func": _abcr.compute_av, "score": "av", "prefer": ["standard"]},
159
+ "sav": {"rule_id": "sav", "func": _abcr.compute_sav, "score": "sav", "prefer": ["standard"]},
160
+ "slav": {"rule_id": "slav", "func": _abcr.compute_slav, "score": "slav", "prefer": ["standard"]},
161
+ "seqpav": {"rule_id": "seqpav", "func": _abcr.compute_seqpav, "score": "sav", "prefer": ["standard"]},
162
+ "seqslav": {"rule_id": "seqslav", "func": _abcr.compute_seqslav, "score": "slav", "prefer": ["standard"]},
163
+ "seqphragmen": {"rule_id": "seqphragmen", "func": _abcr.compute_seqphragmen, "score": "sav", "prefer": ["float-fractions", "standard-fractions", "gmpy2-fractions"]},
164
+ "seqcc": {"rule_id": "seqcc", "func": _abcr.compute_seqcc, "score": "sav", "prefer": ["standard"]},
165
+ "revseqpav": {"rule_id": "revseqpav", "func": _abcr.compute_revseqpav, "score": "sav", "prefer": ["standard"]},
166
+ "phragmen_enestroem":{"rule_id": "phragmen-enestroem","func": _abcr.compute_phragmen_enestroem,"score": "sav", "prefer": ["standard"]},
167
+ "equal_shares": {"rule_id": "equal-shares", "func": _abcr.compute_equal_shares, "score": "sav", "prefer": ["float-fractions", "standard-fractions", "gmpy2-fractions"]},
168
+ "rule_x": {"rule_id": "rule-x", "func": _abcr.compute_rule_x, "score": "sav", "prefer": ["float-fractions", "standard-fractions", "gmpy2-fractions"]},
169
+ "consensus_rule": {"rule_id": "consensus-rule", "func": _abcr.compute_consensus_rule, "score": "av", "prefer": ["float-fractions", "standard-fractions", "gmpy2-fractions"]},
170
+ "eph": {"rule_id": "eph", "func": _abcr.compute_eph, "score": "sav", "prefer": ["float-fractions", "standard-fractions", "gmpy2-fractions"]},
171
+ "rsd": {"rule_id": "rsd", "func": _abcr.compute_rsd, "score": "av", "prefer": ["standard"]},
172
+ "pav": {"rule_id": "pav", "func": _abcr.compute_pav, "score": "sav", "prefer": ["branch-and-bound", "brute-force", "mip-cbc"]},
173
+ "cc": {"rule_id": "cc", "func": _abcr.compute_cc, "score": "sav", "prefer": ["branch-and-bound", "brute-force", "mip-cbc", "ortools-cp"]},
174
+ "monroe": {"rule_id": "monroe", "func": _abcr.compute_monroe, "score": "sav", "prefer": ["brute-force", "mip-cbc", "ortools-cp"]},
175
+ "greedy_monroe": {"rule_id": "greedy-monroe", "func": _abcr.compute_greedy_monroe, "score": "sav", "prefer": ["standard"]},
176
+ "lexcc": {"rule_id": "lexcc", "func": _abcr.compute_lexcc, "score": "sav", "prefer": ["brute-force"]},
177
+ "lexminimaxav": {"rule_id": "lexminimaxav", "func": _abcr.compute_lexminimaxav, "score": "av", "prefer": ["brute-force"]},
178
+ "maximin_support": {"rule_id": "maximin-support", "func": _abcr.compute_maximin_support, "score": "sav", "prefer": ["mip-cbc"]},
179
+ "minimaxphragmen": {"rule_id": "minimaxphragmen", "func": _abcr.compute_minimaxphragmen, "score": "sav", "prefer": ["mip-cbc"]},
180
+ }
181
+
182
+
183
+ def make_abcvoter(rule_id: str, top_k: int, **kwargs) -> ABCVotingRule:
184
+ """
185
+ Factory for creating an ABCVotingRule from the SAFE_RULES catalog.
186
+
187
+ Example
188
+ -------
189
+ >>> agg = make_abcvoter("seqpav", top_k=10)
190
+ >>> agg.fit(R_binary)
191
+ """
192
+ rid = rule_id.lower()
193
+ if rid not in SAFE_RULES:
194
+ raise ValueError(
195
+ f"Rule '{rule_id}' not available. "
196
+ f"Available: {list(SAFE_RULES.keys())}"
197
+ )
198
+ spec = SAFE_RULES[rid]
199
+ prefer_algorithms = list(spec.get("prefer", []))
200
+ if "prefer_algorithms" in kwargs:
201
+ prefer_algorithms = list(kwargs.pop("prefer_algorithms")) + prefer_algorithms
202
+ return ABCVotingRule(
203
+ top_k=top_k,
204
+ compute_func=spec["func"],
205
+ rule_id=spec.get("rule_id"),
206
+ score_mode=spec["score"],
207
+ prefer_algorithms=prefer_algorithms,
208
+ **kwargs,
209
+ )
210
+
211
+
212
+ class ABCVoteAggregator:
213
+ """
214
+ Adapter that wraps ABCVotingRule to the simple aggregator interface:
215
+ - aggregate(boot_results, feature_names) -> DataFrame ['feature','score'] sorted desc.
216
+ - boot_results can be:
217
+ * iterable of subsets (names or indices), or
218
+ * iterable of 1D binary masks (len = n_features).
219
+ """
220
+ def __init__(self, rule: str = "seqpav", top_k: int | None = None,
221
+ prefer_algorithms: Optional[List[str]] = None, **kwargs):
222
+ self.rule = (rule or "seqpav").lower()
223
+ self.top_k = top_k # if None, we infer a reasonable k from the data
224
+ self.prefer_algorithms = list(prefer_algorithms or [])
225
+ self.kwargs = dict(kwargs)
226
+
227
+ def _to_binary_matrix(self, boot_results, feature_names: List[str]) -> np.ndarray:
228
+ n = len(feature_names)
229
+ results = list(boot_results)
230
+ R = np.zeros((len(results), n), dtype=int)
231
+ for b, res in enumerate(results):
232
+ if isinstance(res, (list, tuple, set)):
233
+ if len(res) == 0:
234
+ continue
235
+ first = next(iter(res))
236
+ if isinstance(first, str):
237
+ idxs = [feature_names.index(f) for f in res if f in feature_names]
238
+ else:
239
+ idxs = list(res)
240
+ R[b, idxs] = 1
241
+ else:
242
+ arr = np.asarray(res)
243
+ if arr.ndim != 1 or arr.shape[0] != n:
244
+ raise ValueError("Each bootstrap result must be a 1D mask with length = n_features.")
245
+ R[b, :] = (arr > 0).astype(int)
246
+ return R
247
+
248
+ def aggregate(self, boot_results, feature_names: List[str]) -> pd.DataFrame:
249
+ feature_names = list(feature_names)
250
+ R = self._to_binary_matrix(boot_results, feature_names)
251
+
252
+ # Infer a reasonable committee size if not provided (median approvals per bootstrap)
253
+ if self.top_k is None:
254
+ per_boot = R.sum(axis=1)
255
+ k = int(np.median(per_boot)) if per_boot.size > 0 else 1
256
+ k = max(1, min(k, len(feature_names)))
257
+ else:
258
+ k = int(self.top_k)
259
+
260
+ voter = make_abcvoter(self.rule, top_k=k, prefer_algorithms=self.prefer_algorithms, **self.kwargs)
261
+ scores = voter.aggregate(R) # np.ndarray length = n_features
262
+
263
+ df = pd.DataFrame({"feature": feature_names, "score": scores})
264
+ return df.sort_values("score", ascending=False, ignore_index=True)
265
+
266
+ # ensure symbol is exported
267
+ try:
268
+ __all__.append("ABCVoteAggregator") # type: ignore[name-defined]
269
+ except Exception:
270
+ __all__ = ["ABCVoteAggregator"]
@@ -0,0 +1,170 @@
1
+ # -*- coding: utf-8 -*-
2
+ from __future__ import annotations
3
+ from abc import ABC, abstractmethod
4
+ from typing import Optional
5
+ import numpy as np
6
+
7
+
8
+ class BaseAggregator(ABC):
9
+ """
10
+ Base class for aggregating feature selection results.
11
+
12
+ Contract:
13
+ - aggregate(results) -> 1D array (n_features,)
14
+ - Input 'results' must be 2D: (n_bootstraps, n_features)
15
+ - By default (ScoreAggregators), larger values are better.
16
+ - Subclasses may override fit() to invert the criterion (e.g., RankAggregator).
17
+ """
18
+
19
+ def __init__(self, top_k: Optional[int] = None):
20
+ self.top_k: Optional[int] = top_k
21
+ self.final_ranking_: Optional[np.ndarray] = None # indices ordered best→worst
22
+ self.selected_features_: Optional[np.ndarray] = None # binary mask (1 selected)
23
+ self._agg_result: Optional[np.ndarray] = None # 1D aggregated vector (scores or ranks)
24
+ self.n_features_: Optional[int] = None
25
+
26
+
27
+ @staticmethod
28
+ def _validate_results(results: np.ndarray) -> np.ndarray:
29
+ """Ensure results is a numeric 2D array (n_bootstraps, n_features)."""
30
+ if results is None:
31
+ raise ValueError("`results` cannot be None.")
32
+ arr = np.asarray(results)
33
+ if arr.ndim != 2:
34
+ raise ValueError(f"`results` must be 2D (n_bootstraps, n_features). Got shape={arr.shape}.")
35
+ if arr.size == 0:
36
+ raise ValueError("`results` is empty.")
37
+ arr = np.nan_to_num(arr.astype(float), nan=0.0, posinf=0.0, neginf=0.0)
38
+ return arr
39
+
40
+ @staticmethod
41
+ def _clip_k(top_k: Optional[int], n_features: int) -> Optional[int]:
42
+ if top_k is None:
43
+ return None
44
+ k = int(top_k)
45
+ if k <= 0:
46
+ raise ValueError(f"`top_k` must be positive. Got {top_k}.")
47
+ return int(min(k, n_features))
48
+
49
+ def set_top_k(self, top_k: Optional[int]) -> "BaseAggregator":
50
+ """Set or update top_k after construction."""
51
+ self.top_k = top_k
52
+ return self
53
+
54
+
55
+ @abstractmethod
56
+ def aggregate(self, results: np.ndarray) -> np.ndarray:
57
+ """
58
+ Parameters
59
+ ----------
60
+ results : np.ndarray of shape (n_bootstraps, n_features)
61
+ Per-bootstrap results to aggregate.
62
+
63
+ Returns
64
+ -------
65
+ aggregated : np.ndarray of shape (n_features,)
66
+ Aggregated vector (scores or ranks).
67
+ """
68
+ raise NotImplementedError
69
+
70
+
71
+ def fit(self, results: np.ndarray):
72
+ """Default fit: assumes higher aggregated values are better (score-like)."""
73
+ res = self._validate_results(results)
74
+ n_boot, n_feat = res.shape
75
+ self.n_features_ = n_feat
76
+
77
+ agg_result = np.asarray(self.aggregate(res)).ravel()
78
+ if agg_result.shape[0] != n_feat:
79
+ raise ValueError(f"aggregate() must return shape (n_features,), got {agg_result.shape}.")
80
+ self._agg_result = agg_result
81
+
82
+ k = self._clip_k(self.top_k, n_feat)
83
+ sorted_indices = np.argsort(-agg_result, kind="mergesort")
84
+ self.final_ranking_ = sorted_indices[:k] if k is not None else sorted_indices
85
+
86
+ mask = np.zeros(n_feat, dtype=int)
87
+ mask[self.final_ranking_] = 1
88
+ self.selected_features_ = mask
89
+ return self
90
+
91
+
92
+ @property
93
+ def scores_(self) -> Optional[np.ndarray]:
94
+ """Return raw aggregated values (scores or ranks depending on subclass)."""
95
+ return self._agg_result
96
+
97
+ @property
98
+ def rank_(self) -> Optional[np.ndarray]:
99
+ """
100
+ Return indices of features ordered best→worst.
101
+ This is an ordered list of indices, NOT per-feature rank numbers.
102
+ """
103
+ return self.final_ranking_
104
+
105
+ def get_support(self, indices: bool = False):
106
+ """
107
+ Scikit-learn-like support API.
108
+ - If indices=False: returns boolean mask (n_features,)
109
+ - If indices=True : returns selected indices (1D)
110
+ """
111
+ if self.selected_features_ is None:
112
+ raise RuntimeError("Estimator not fitted yet; `selected_features_` is None.")
113
+ mask_bool = np.asarray(self.selected_features_, dtype=bool)
114
+ return np.flatnonzero(mask_bool) if indices else mask_bool
115
+
116
+
117
+ class ScoreAggregator(BaseAggregator):
118
+ """Aggregators where higher score = better (use BaseAggregator.fit)."""
119
+ pass
120
+
121
+
122
+ class RankAggregator(BaseAggregator):
123
+ """
124
+ Aggregators producing per-feature rank values, where LOWER = better.
125
+ aggregate() must return per-feature rank values (e.g., mean rank).
126
+ """
127
+ def fit(self, results: np.ndarray):
128
+ res = self._validate_results(results)
129
+ n_boot, n_feat = res.shape
130
+ self.n_features_ = n_feat
131
+
132
+ rank_values = np.asarray(self.aggregate(res)).ravel()
133
+ if rank_values.shape[0] != n_feat:
134
+ raise ValueError(f"aggregate() must return shape (n_features,), got {rank_values.shape}.")
135
+ self._agg_result = rank_values
136
+
137
+ k = self._clip_k(self.top_k, n_feat)
138
+ sorted_indices = np.argsort(rank_values, kind="mergesort")
139
+ self.final_ranking_ = sorted_indices[:k] if k is not None else sorted_indices
140
+
141
+ mask = np.zeros(n_feat, dtype=int)
142
+ mask[self.final_ranking_] = 1
143
+ self.selected_features_ = mask
144
+ return self
145
+
146
+
147
+ class BinaryAggregator(BaseAggregator):
148
+ """
149
+ Aggregators that produce a binary selection mask (0/1).
150
+ aggregate() must return a 0/1 vector; fit() stores it and builds a stable order.
151
+ """
152
+ def fit(self, results: np.ndarray):
153
+ res = self._validate_results(results)
154
+ n_boot, n_feat = res.shape
155
+ self.n_features_ = n_feat
156
+
157
+ mask = np.asarray(self.aggregate(res)).astype(int).ravel()
158
+ if mask.shape[0] != n_feat:
159
+ raise ValueError(f"aggregate() must return shape (n_features,), got {mask.shape}.")
160
+ mask = (mask > 0).astype(int)
161
+
162
+ self._agg_result = mask
163
+ self.selected_features_ = mask
164
+
165
+ ones = np.where(mask == 1)[0]
166
+ zeros = np.where(mask == 0)[0]
167
+ order = np.concatenate([ones, zeros])
168
+ k = self._clip_k(self.top_k, n_feat)
169
+ self.final_ranking_ = order[:k] if k is not None else order
170
+ return self