measly 0.1.0__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.
measly/__init__.py ADDED
@@ -0,0 +1,43 @@
1
+ """Estimate whether a model is data-limited, and whether more capacity would pay off."""
2
+
3
+ from measly.analysis import (
4
+ DEFAULT_FRACTIONS,
5
+ VALIDATION_WARN,
6
+ Analysis,
7
+ Projection,
8
+ analyse,
9
+ mean_squared_error,
10
+ )
11
+ from measly.fit import (
12
+ POW3,
13
+ POW4,
14
+ CurveEnsemble,
15
+ ScalingLaw,
16
+ fit_scaling_law,
17
+ fit_scaling_laws,
18
+ pow3,
19
+ )
20
+ from measly.interfaces import Model, Score
21
+ from measly.plot import plot
22
+ from measly.sweep import sweep, train_test_split
23
+
24
+ __all__ = [
25
+ "analyse",
26
+ "Analysis",
27
+ "Projection",
28
+ "mean_squared_error",
29
+ "plot",
30
+ "DEFAULT_FRACTIONS",
31
+ "VALIDATION_WARN",
32
+ "Model",
33
+ "Score",
34
+ "train_test_split",
35
+ "sweep",
36
+ "pow3",
37
+ "ScalingLaw",
38
+ "POW3",
39
+ "POW4",
40
+ "CurveEnsemble",
41
+ "fit_scaling_law",
42
+ "fit_scaling_laws",
43
+ ]
measly/analysis.py ADDED
@@ -0,0 +1,287 @@
1
+ from collections.abc import Mapping, Sequence
2
+ from dataclasses import dataclass
3
+
4
+ import numpy as np
5
+ import xarray as xr
6
+ from numpy.typing import NDArray
7
+
8
+ from measly.fit import POW4, CurveEnsemble, ScalingLaw, fit_scaling_law, fit_scaling_laws
9
+ from measly.interfaces import Model, Score
10
+ from measly.sweep import sweep
11
+
12
+ __all__ = [
13
+ "mean_squared_error",
14
+ "DEFAULT_FRACTIONS",
15
+ "VALIDATION_WARN",
16
+ "Projection",
17
+ "Analysis",
18
+ "analyse",
19
+ ]
20
+
21
+ DEFAULT_FRACTIONS: tuple[float, ...] = tuple(
22
+ round(f, 4) for f in np.geomspace(0.1, 1.0, 8)
23
+ )
24
+
25
+
26
+ VALIDATION_WARN = 0.15
27
+ """Above this held-back error, a projection MUST NOT be relied upon."""
28
+
29
+
30
+ def mean_squared_error(predicted, true) -> float:
31
+ """The default `Score`. Regression only. Supply your own otherwise."""
32
+ return float(np.mean((np.asarray(predicted) - np.asarray(true)) ** 2))
33
+
34
+
35
+ def _score_name(score: Score) -> str:
36
+ """A label for the report. Falls back to the type for a callable object."""
37
+ return getattr(score, "__name__", None) or type(score).__name__
38
+
39
+
40
+ def _unwrap(values: NDArray, scalar: bool) -> float | NDArray:
41
+ """Give back a float for a scalar `factor`, an array for several."""
42
+ return float(values[0]) if scalar else values
43
+
44
+
45
+ @dataclass(frozen=True)
46
+ class Projection:
47
+ """What a model is expected to score on a larger training set.
48
+
49
+ Every field is a float when `project` is given one factor, and an array
50
+ aligned with the factors when it is given several.
51
+ """
52
+
53
+ loss: float | NDArray
54
+ low: float | NDArray
55
+ high: float | NDArray
56
+ gain: float | NDArray
57
+ gain_low: float | NDArray
58
+ gain_high: float | NDArray
59
+
60
+
61
+ @dataclass(frozen=True)
62
+ class Analysis:
63
+ """The measured curve, the fitted curves, and whether to believe them."""
64
+
65
+ results: xr.DataArray
66
+ curves: dict[str, CurveEnsemble]
67
+ validation: dict[str, float]
68
+ law: ScalingLaw
69
+ score: Score
70
+ n_train: int
71
+
72
+ @property
73
+ def models(self) -> list[str]:
74
+ return list(self.curves)
75
+
76
+ def measured(self) -> dict[str, float]:
77
+ """Mean score at the full training pool, straight from the sweep."""
78
+ largest = float(self.results.fraction.max())
79
+ return {
80
+ name: float(self.results.sel(model=name, fraction=largest).mean("draw"))
81
+ for name in self.models
82
+ }
83
+
84
+ def project(self, factor: float = 2.0, interval: float = 0.9
85
+ ) -> dict[str, Projection]:
86
+ """Loss and gain at `factor` times the current training pool.
87
+
88
+ `factor` MAY be a sequence. One call then gives the whole
89
+ extrapolation curve, which is what a plot needs.
90
+
91
+ The interval on `loss` is a percentile range across the fitted curves.
92
+ It is dominated by test-set noise: every point in a draw shares one
93
+ test set, so that draw's curve is shifted bodily. It is conservative,
94
+ covering 98% or more at a nominal 90% against simulated ground truth.
95
+
96
+ The interval on `gain` is far tighter, because each member is paired
97
+ with itself and the shared offset cancels. `gain` is therefore not
98
+ `measured()` minus `loss`; it is what the curve itself says it climbs.
99
+ """
100
+ tail = (1.0 - interval) / 2.0
101
+ quantiles = [tail, 1.0 - tail]
102
+
103
+ scalar = np.ndim(factor) == 0
104
+ factors = np.atleast_1d(np.asarray(factor, dtype=float))
105
+
106
+ projections = {}
107
+ for name, ensemble in self.curves.items():
108
+ # (len(factors), n_members). Quantiles MUST reduce the member axis
109
+ # only; flattening both would mix factors into one distribution.
110
+ future = ensemble.predict(factors)
111
+ # Pair each member against its own value at the measured size.
112
+ # Both carry that draw's test-set offset, so it cancels; comparing
113
+ # against a mean over draws leaves it in and swamps the gain.
114
+ gains = ensemble.predict(np.ones(1)) - future
115
+ lo, hi = np.quantile(future, quantiles, axis=-1)
116
+ gain_lo, gain_hi = np.quantile(gains, quantiles, axis=-1)
117
+ projections[name] = Projection(
118
+ loss=_unwrap(np.median(future, axis=-1), scalar),
119
+ low=_unwrap(lo, scalar),
120
+ high=_unwrap(hi, scalar),
121
+ gain=_unwrap(np.median(gains, axis=-1), scalar),
122
+ gain_low=_unwrap(gain_lo, scalar),
123
+ gain_high=_unwrap(gain_hi, scalar),
124
+ )
125
+ return projections
126
+
127
+ def summary(self, factor: float = 2.0) -> str:
128
+ now = self.measured()
129
+ projected = self.project(factor)
130
+ target = round(self.n_train * factor)
131
+
132
+ lines = [
133
+ f"measly: {len(self.models)} model(s), {self.n_train} training examples, "
134
+ f"{self.results.sizes['draw']} draws",
135
+ f" {self.law!r}, {_score_name(self.score)} (lower is better)",
136
+ "",
137
+ ]
138
+ for name in self.models:
139
+ p = projected[name]
140
+
141
+ error = self.validation.get(name, float("nan"))
142
+ if np.isnan(error):
143
+ check = "not run"
144
+ elif error > VALIDATION_WARN:
145
+ check = (f"{error:.1%} error — this curve does not predict its own "
146
+ "measured sizes, so do not rely on the projection")
147
+ else:
148
+ check = f"{error:.1%} error"
149
+
150
+ rows = [
151
+ (f"now ({self.n_train} examples)", f"{now[name]:.4f}"),
152
+ (f"at {target} examples",
153
+ f"{p.loss:.4f} [{p.low:.4f}, {p.high:.4f}]"),
154
+ ("gain from getting there",
155
+ f"{p.gain:+.4f} [{p.gain_low:+.4f}, {p.gain_high:+.4f}]"),
156
+ ("held-back check", check),
157
+ ]
158
+ width = max(len(label) for label, _ in rows)
159
+
160
+ lines.append(f" {name}")
161
+ lines += [f" {label:<{width}} {value}" for label, value in rows]
162
+ lines.append("")
163
+
164
+ return "\n".join(lines).rstrip()
165
+
166
+ def __repr__(self) -> str:
167
+ worst = max((v for v in self.validation.values() if not np.isnan(v)),
168
+ default=float("nan"))
169
+ checked = "unchecked" if np.isnan(worst) else f"worst held-back {worst:.1%}"
170
+ return (f"<Analysis {len(self.models)} model(s), n_train={self.n_train}, "
171
+ f"{self.law!r}, {checked}>")
172
+
173
+
174
+ def _validate(results: xr.DataArray, fractions: Sequence[float], hold_back: int,
175
+ law: ScalingLaw, floor: float) -> dict[str, float]:
176
+ """Refit on the smaller sizes, then predict the largest measured ones.
177
+
178
+ Returns the worst relative error per model. Returns nan when too few points
179
+ remain to fit the law once some are held back.
180
+ """
181
+ names = [str(n) for n in np.atleast_1d(results.model.values)]
182
+ kept, held = list(fractions[:-hold_back]), list(fractions[-hold_back:])
183
+
184
+ n_params = law.func.__code__.co_argcount - 1
185
+ if hold_back <= 0 or len(kept) <= n_params:
186
+ return {name: float("nan") for name in names}
187
+
188
+ errors = {}
189
+ for name in names:
190
+ ensemble = fit_scaling_law(
191
+ results.sel(model=name, fraction=kept), law=law, floor=floor
192
+ )
193
+ if len(ensemble) == 0:
194
+ errors[name] = float("nan")
195
+ continue
196
+
197
+ worst = 0.0
198
+ for fraction in held:
199
+ actual = float(results.sel(model=name, fraction=fraction).mean("draw"))
200
+ predicted = float(np.median(np.ravel(ensemble.predict(fraction))))
201
+ if actual != 0.0:
202
+ worst = max(worst, abs(predicted - actual) / abs(actual))
203
+ errors[name] = worst
204
+
205
+ return errors
206
+
207
+
208
+ def analyse(
209
+ models: Mapping[str, Model] | Sequence[Model],
210
+ X,
211
+ y,
212
+ score: Score = mean_squared_error,
213
+ *,
214
+ fractions: Sequence[float] = DEFAULT_FRACTIONS,
215
+ n_draws: int = 100,
216
+ test_fraction: float = 0.25,
217
+ law: ScalingLaw = POW4,
218
+ floor: float = 0.0,
219
+ rng: int | np.random.Generator | None = None,
220
+ hold_back: int = 2,
221
+ n_jobs: int = 1,
222
+ ) -> Analysis:
223
+ """Measure a learning curve for each model and project it forward.
224
+
225
+ `rng` accepts a seed, a `Generator`, or nothing. Pass a seed for a
226
+ reproducible run. The split and the subsets both depend on it.
227
+
228
+ `hold_back` sets how many of the largest measured sizes are withheld from a
229
+ second fit, used to check the projection. Set it to 0 to skip the check.
230
+ You then have no evidence the form suits your data.
231
+
232
+ `models` MAY be a `{name: model}` mapping. Name them whenever `repr` is
233
+ unhelpful, as it is for an sklearn `Pipeline`.
234
+
235
+ `n_jobs` parallelises the sweep without changing results. See `sweep`.
236
+ """
237
+ generator = np.random.default_rng(rng)
238
+ fractions = list(fractions)
239
+
240
+ results = sweep(
241
+ models=models,
242
+ X=X,
243
+ y=y,
244
+ fractions=fractions,
245
+ rng=generator,
246
+ score=score,
247
+ n_draws=n_draws,
248
+ test_fraction=test_fraction,
249
+ n_jobs=n_jobs,
250
+ )
251
+
252
+ curves = fit_scaling_laws(results, law=law, floor=floor)
253
+ validation = _validate(results, fractions, hold_back, law, floor)
254
+ n_train = int(round(len(y) * (1.0 - test_fraction) * max(fractions)))
255
+
256
+ return Analysis(results=results, curves=curves, validation=validation,
257
+ law=law, score=score, n_train=n_train)
258
+
259
+
260
+ if __name__ == "__main__":
261
+ # uv run python src/measly/analysis.py
262
+ # The `-m` form warns: __init__ imports this module, so runpy loads it twice.
263
+ class _Ridge:
264
+ """A model in the fit/predict convention. Penalty sets its capacity."""
265
+
266
+ def __init__(self, penalty: float) -> None:
267
+ self.penalty, self.weights = penalty, None
268
+
269
+ def __repr__(self) -> str:
270
+ return f"ridge({self.penalty:g})"
271
+
272
+ def fit(self, X, y) -> None:
273
+ gram = X.T @ X + self.penalty * np.eye(X.shape[1])
274
+ self.weights = np.linalg.solve(gram, X.T @ y)
275
+
276
+ def predict(self, X):
277
+ return X @ self.weights
278
+
279
+ _rng = np.random.default_rng(0)
280
+ _X = _rng.normal(size=(2000, 12))
281
+ _y = _X @ _rng.normal(size=12) + _rng.normal(0.0, 1.0, 2000)
282
+ _models = [_Ridge(1.0), _Ridge(200.0)]
283
+
284
+ # One call does the sweep, the fits and the hold-back check.
285
+ _result = analyse(_models, _X, _y, rng=0, score=mean_squared_error)
286
+ print(_result.summary())
287
+
measly/fit.py ADDED
@@ -0,0 +1,145 @@
1
+ """Fit a scaling law to a sweep's measured losses.
2
+
3
+ One curve per draw. Draws are independent, so a single curve means nothing.
4
+ Read percentiles across the ensemble.
5
+ """
6
+
7
+ from dataclasses import dataclass, field
8
+ from typing import Callable
9
+
10
+ import numpy as np
11
+ import xarray as xr
12
+ from numpy.typing import ArrayLike, NDArray
13
+ from scipy.optimize import curve_fit
14
+
15
+
16
+ __all__ = [
17
+ "pow3",
18
+ "ScalingLaw",
19
+ "POW3",
20
+ "POW4",
21
+ "CurveEnsemble",
22
+ "fit_scaling_law",
23
+ "fit_scaling_laws",
24
+ ]
25
+
26
+ # Below 0, loss would rise with data. Above 5 is not physical.
27
+ _ALPHA_BOUNDS = (0.01, 5.0)
28
+
29
+
30
+ def pow3(n, lower_bound, A, alpha):
31
+ """L(n) = L_inf + A * n^(-alpha). Called `pow3` in the literature.
32
+
33
+ `n` MAY be counts or fractions of the pool. `lower_bound` and `alpha` are
34
+ identical either way. Rescaling n is absorbed entirely by `A`.
35
+ """
36
+ return lower_bound + A * np.asarray(n, dtype=float) ** (-alpha)
37
+
38
+
39
+ @dataclass(frozen=True)
40
+ class ScalingLaw:
41
+ """A functional form plus the starting point and bounds it needs.
42
+
43
+ These MUST travel together. Forms differ in parameter count, and in which
44
+ parameter is the asymptote that `floor` constrains. A generic guess would
45
+ constrain the wrong one. `guess(y, floor)` returns `(p0, (lower, upper))`.
46
+ """
47
+
48
+ name: str
49
+ func: Callable[..., NDArray]
50
+ guess: Callable[[NDArray, float], tuple[list[float], tuple]]
51
+
52
+ def __call__(self, x, *params) -> NDArray:
53
+ return self.func(x, *params)
54
+
55
+ def __repr__(self) -> str:
56
+ return self.name
57
+
58
+
59
+ def _pow4(n, lower_bound, A, alpha, shift):
60
+ """L(n) = L_inf + A * (n + shift)^(-alpha). Called `pow4`.
61
+
62
+ `pow3` is this with `shift` at zero. So `pow4` never fits worse.
63
+
64
+ `shift` separates where the curve bends from how low it ends up. Under
65
+ `pow3` one parameter sets both. A curve still descending at the largest
66
+ measured size can then only be fitted by driving `L_inf` to zero.
67
+ """
68
+ return lower_bound + A * (np.asarray(n, dtype=float) + shift) ** (-alpha)
69
+
70
+
71
+ def _pow3_guess(y: NDArray, floor: float) -> tuple[list[float], tuple]:
72
+ start = max(float(y.min()) * 0.9, floor)
73
+ p0 = [start, max(float(y.max() - y.min()), 1e-6), 1.0]
74
+ return p0, ((floor, 0.0, _ALPHA_BOUNDS[0]), (np.inf, np.inf, _ALPHA_BOUNDS[1]))
75
+
76
+
77
+ def _pow4_guess(y: NDArray, floor: float) -> tuple[list[float], tuple]:
78
+ p0, (lower, upper) = _pow3_guess(y, floor)
79
+ return p0 + [0.01], (lower + (0.0,), upper + (10.0,))
80
+
81
+
82
+ POW3 = ScalingLaw("pow3", pow3, _pow3_guess)
83
+ POW4 = ScalingLaw("pow4", _pow4, _pow4_guess)
84
+
85
+
86
+ @dataclass
87
+ class CurveEnsemble:
88
+ """Fitted curves from one model's draws. Read them by evaluating all."""
89
+
90
+ curves: list[Callable] = field(default_factory=list)
91
+ n_failed: int = 0 # draws whose fit did not converge
92
+
93
+ def __len__(self) -> int:
94
+ return len(self.curves)
95
+
96
+ def predict(self, x: ArrayLike) -> NDArray:
97
+ """Evaluate every member at every x.
98
+
99
+ Returns `(len(x), n_members)`. Take quantiles across members for an
100
+ interval. A single member MUST NOT be read on its own.
101
+ """
102
+ return np.stack([curve(x) for curve in self.curves], axis=-1)
103
+
104
+
105
+ def fit_scaling_law(
106
+ model_results: xr.DataArray, law: ScalingLaw = POW4, floor: float = 0.0
107
+ ) -> CurveEnsemble:
108
+ """Fit one curve per draw for a single model.
109
+
110
+ `model_results` has dims `(fraction, draw)`. `floor` raises the lower bound
111
+ on the asymptote. Pass a measured noise floor if you have one.
112
+
113
+ `law` picks the form. `POW4` is the default. On curves that have not yet
114
+ bent it extrapolates far better than `POW3`. `POW3` in that regime loses to
115
+ assuming no further improvement at all. `POW3` is what the literature
116
+ usually quotes.
117
+ """
118
+ ensemble = CurveEnsemble()
119
+
120
+ x = np.asarray(model_results.fraction.values, dtype=float)
121
+
122
+ for y in model_results.transpose("draw", "fraction").values:
123
+ p0, bounds = law.guess(y, floor)
124
+
125
+ try:
126
+ popt, _ = curve_fit(law.func, x, y, p0=p0, bounds=bounds, maxfev=10_000)
127
+ except (RuntimeError, ValueError):
128
+ ensemble.n_failed += 1
129
+ continue
130
+
131
+ # popt MUST be bound as a default argument. A closure over the loop
132
+ # variable leaves every curve holding the last draw's parameters.
133
+ ensemble.curves.append(lambda t, p=popt: law.func(t, *p))
134
+
135
+ return ensemble
136
+
137
+
138
+ def fit_scaling_laws(
139
+ results: xr.DataArray, law: ScalingLaw = POW4, floor: float = 0.0
140
+ ) -> dict[str, CurveEnsemble]:
141
+ """Fit every model in a sweep, keyed by its label on the `model` axis."""
142
+ return {
143
+ str(name): fit_scaling_law(results.sel(model=name), law=law, floor=floor)
144
+ for name in results.model.values
145
+ }
measly/interfaces.py ADDED
@@ -0,0 +1,31 @@
1
+ """The contracts a user implements to run measly on their own problem."""
2
+
3
+ from typing import Any, Callable, Protocol, Sequence
4
+
5
+ __all__ = ["Score", "Model"]
6
+
7
+
8
+ type Score = Callable[[Sequence[float], Sequence[float]], float]
9
+ """Reduce predictions and true targets to one number. Lower MUST be better.
10
+
11
+ It SHOULD be a mean over per-point terms: MSE, MAE, error rate. Set-level
12
+ statistics MUST NOT be used. F1, AUC and Spearman do not decompose per point,
13
+ so the power-law form does not describe them. The fit still returns a number.
14
+ """
15
+
16
+
17
+ class Model(Protocol):
18
+ """Anything with scikit-learn's fit/predict convention.
19
+
20
+ `fit` MUST leave the model fully retrained. `sweep` refits one instance at
21
+ every sample size. A model that carries state between calls, such as an
22
+ sklearn estimator with `warm_start=True`, reports the same score at every
23
+ size. That plots as a flat curve and raises nothing.
24
+
25
+ `predict` MUST return one number per row of `X`, in order. Returns are
26
+ `Any`: sklearn's `fit` gives back `self`, and `predict` an ndarray.
27
+ """
28
+
29
+ def fit(self, X, y) -> Any: ...
30
+
31
+ def predict(self, X) -> Any: ...
measly/plot.py ADDED
@@ -0,0 +1,81 @@
1
+ """One figure for an `Analysis`: what was measured, and what follows from it."""
2
+
3
+ import matplotlib.pyplot as plt
4
+ import numpy as np
5
+ from scipy.ndimage import gaussian_filter1d
6
+ from matplotlib.ticker import FuncFormatter, LogLocator, NullLocator
7
+
8
+ from measly.analysis import Analysis, _score_name
9
+
10
+ __all__ = ["plot"]
11
+
12
+ PALETTE = ("#2563eb", "#ea580c", "#059669", "#7c3aed", "#db2777")
13
+ GUIDE = "#94a3b8"
14
+
15
+
16
+ def plot(analysis: Analysis, factor: float = 2.0, interval: float = 0.9, ax=None):
17
+ """Measured points, the fitted band, and the projection out to `factor`.
18
+
19
+ Points are the mean over draws, with bars at one standard deviation. The
20
+ band is the percentile range across fitted curves, so it widens right of
21
+ the dashed line, where nothing was measured.
22
+
23
+ Returns the `Axes`, so several analyses MAY share a figure.
24
+ """
25
+ if ax is None:
26
+ _, ax = plt.subplots(figsize=(7.5, 4.5))
27
+
28
+ measured = np.asarray(analysis.results.fraction.values, dtype=float)
29
+ grid = np.geomspace(measured.min(), max(factor, measured.max()), 200)
30
+ projected = analysis.project(grid, interval)
31
+
32
+ for index, name in enumerate(analysis.models):
33
+ scores = analysis.results.sel(model=name)
34
+ band = projected[name]
35
+ colour = PALETTE[index % len(PALETTE)]
36
+
37
+ error = analysis.validation.get(name, float("nan"))
38
+ label = name if np.isnan(error) else f"{name} held-back {error:.1%}"
39
+
40
+ ax.fill_between(grid * analysis.n_train, band.low, band.high,
41
+ color=colour, alpha=0.12, lw=0)
42
+ for edge in (band.low, band.high):
43
+ ax.plot(grid * analysis.n_train, edge, color=colour, alpha=0.35, lw=0.8)
44
+ # The pointwise median of a finite ensemble kinks where members cross.
45
+ # Smoothing is cosmetic. `project` and `summary` report the raw median.
46
+ centre = gaussian_filter1d(band.loss, sigma=4, mode="nearest")
47
+ ax.plot(grid * analysis.n_train, centre, color=colour, lw=2,
48
+ label=label, zorder=4)
49
+ ax.errorbar(measured * analysis.n_train, scores.mean("draw"),
50
+ yerr=scores.std("draw", ddof=1), fmt="o", ms=5,
51
+ color=colour, elinewidth=1.1, capsize=0, zorder=5)
52
+
53
+ ax.axvline(analysis.n_train, color=GUIDE, ls=(0, (2, 3)), lw=1, zorder=1)
54
+ for text, offset, side in (("measured", -5, "right"), ("projected", 5, "left")):
55
+ ax.annotate(text, xy=(analysis.n_train, 0.0),
56
+ xycoords=("data", "axes fraction"), xytext=(offset, 7),
57
+ textcoords="offset points", ha=side, va="bottom",
58
+ color=GUIDE, fontsize=8)
59
+
60
+ ax.set_xscale("log")
61
+ ax.set_xlim(min(grid) * analysis.n_train / 1.05, max(grid) * analysis.n_train * 1.05)
62
+ ax.xaxis.set_major_locator(LogLocator(subs=(1, 2, 3, 5)))
63
+ ax.xaxis.set_minor_locator(NullLocator())
64
+ ax.xaxis.set_major_formatter(FuncFormatter(lambda v, _: f"{v:,.0f}"))
65
+
66
+ ax.grid(axis="y", color="#e5e7eb", lw=0.8)
67
+ ax.set_axisbelow(True)
68
+ for side in ("top", "right"):
69
+ ax.spines[side].set_visible(False)
70
+ for side in ("left", "bottom"):
71
+ ax.spines[side].set_color(GUIDE)
72
+ ax.tick_params(colors="#475569", length=4)
73
+
74
+ ax.set_xlabel("training examples", color="#334155")
75
+ ax.set_ylabel(_score_name(analysis.score), color="#334155")
76
+ ax.set_title(f"{analysis.law!r} fitted to {analysis.n_train} examples, "
77
+ f"projected to {round(analysis.n_train * factor):,}",
78
+ loc="left", color="#1e293b", fontsize=11, pad=14)
79
+ ax.legend(frameon=False, fontsize=9, loc="upper right",
80
+ handlelength=1.6, borderaxespad=0)
81
+ return ax
measly/sweep.py ADDED
@@ -0,0 +1,97 @@
1
+ """Measure performance across a grid of training-set sizes."""
2
+
3
+ import os
4
+ from collections.abc import Mapping, Sequence
5
+ from concurrent.futures import ProcessPoolExecutor
6
+
7
+ import numpy as np
8
+ import xarray as xr
9
+
10
+ from measly.interfaces import Score, Model
11
+
12
+ __all__ = ["train_test_split", "sweep"]
13
+
14
+
15
+ def train_test_split(X, y, test_fraction: float, rng: np.random.Generator) -> tuple:
16
+
17
+ order = rng.permutation(len(X))
18
+ n_test = int(round(len(y) * test_fraction))
19
+ test, train = order[:n_test], order[n_test:]
20
+
21
+ return X[train], X[test], y[train], y[test]
22
+
23
+
24
+ def _score_draw(task) -> np.ndarray:
25
+ """One draw's scores, `(fraction, model)`. Top level so a pool can pickle it."""
26
+ models, score, train_X, train_y, test_X, test_y, subsets = task
27
+ scores = np.empty((len(subsets), len(models)))
28
+ for m, model in enumerate(models):
29
+ for f, indices in enumerate(subsets):
30
+ model.fit(train_X[indices], train_y[indices])
31
+ scores[f, m] = score(model.predict(test_X), test_y)
32
+ return scores
33
+
34
+
35
+ def sweep(
36
+ models: Mapping[str, Model] | Sequence[Model],
37
+ X,
38
+ y,
39
+ fractions: list[float],
40
+ rng: np.random.Generator,
41
+ score: Score,
42
+ n_draws: int = 1000,
43
+ test_fraction: float = 0.25,
44
+ n_jobs: int = 1,
45
+ ) -> xr.DataArray:
46
+ """Score every model at every fraction of the training pool, `n_draws` times.
47
+
48
+ Only the training pool is downsampled. The test set is fixed within a draw
49
+ and redrawn between draws. Models MUST be scored on held-out rows. Scoring
50
+ on the training rows makes the curve rise with n rather than fall.
51
+
52
+ Subsets are drawn without replacement. Fraction 1.0 still varies across
53
+ draws, because each draw resplits the pool. A bootstrap is not needed for
54
+ that spread and biases the level up by repeating rows.
55
+
56
+ `models` MAY be a `{name: model}` mapping, which then keys the `model`
57
+ axis. Otherwise `repr` does, and an sklearn `Pipeline`'s embeds an address.
58
+ Those keys MUST be distinct, or one model would overwrite another.
59
+
60
+ `n_jobs` > 1 scores draws in that many processes, -1 in all cores. Results
61
+ do not change. `score` and the models MUST be picklable, so no lambdas.
62
+ Set `OMP_NUM_THREADS=1` or the workers' BLAS threads will contend.
63
+ """
64
+
65
+ named = (dict(models) if isinstance(models, Mapping)
66
+ else {repr(model): model for model in models})
67
+ if len(named) != len(models):
68
+ raise ValueError(
69
+ "models must have distinct repr(). The results array keys its model "
70
+ "axis by repr, so duplicates would overwrite each other. Pass a "
71
+ "{name: model} mapping to name them."
72
+ )
73
+
74
+ def tasks():
75
+ for _ in range(n_draws):
76
+ train_X, test_X, train_y, test_y = train_test_split(X, y, test_fraction, rng)
77
+ subsets = [
78
+ rng.choice(len(train_X), size=round(len(train_X) * fraction), replace=False)
79
+ for fraction in fractions
80
+ ]
81
+ yield list(named.values()), score, train_X, train_y, test_X, test_y, subsets
82
+
83
+ if n_jobs == 1:
84
+ per_draw = map(_score_draw, tasks())
85
+ else:
86
+ # ponytail: holds all draws' data at once. Ship indices if the pool is large.
87
+ workers = (os.cpu_count() or 1) if n_jobs == -1 else n_jobs
88
+ with ProcessPoolExecutor(workers) as pool:
89
+ per_draw = list(pool.map(_score_draw, tasks(),
90
+ chunksize=max(1, n_draws // (4 * workers))))
91
+ scores = np.stack(list(per_draw), axis=-1)
92
+
93
+ return xr.DataArray(
94
+ scores,
95
+ dims=("fraction", "model", "draw"),
96
+ coords={"fraction": fractions, "model": list(named), "draw": np.arange(n_draws)},
97
+ )
@@ -0,0 +1,122 @@
1
+ Metadata-Version: 2.5
2
+ Name: measly
3
+ Version: 0.1.0
4
+ Summary: Estimate whether a model is data-limited, and whether more capacity would pay off.
5
+ Project-URL: Homepage, https://github.com/tremgan/measly
6
+ Project-URL: Source, https://github.com/tremgan/measly
7
+ Project-URL: Issues, https://github.com/tremgan/measly/issues
8
+ Author-email: Remi Tregan <tregan.remi@gmail.com>
9
+ License-Expression: MIT
10
+ License-File: LICENSE
11
+ Keywords: active-learning,learning-curve,sample-size,scaling-law
12
+ Classifier: Development Status :: 3 - Alpha
13
+ Classifier: Intended Audience :: Science/Research
14
+ Classifier: Programming Language :: Python :: 3.12
15
+ Classifier: Programming Language :: Python :: 3.13
16
+ Classifier: Topic :: Scientific/Engineering
17
+ Requires-Python: >=3.12
18
+ Requires-Dist: matplotlib>=3.8
19
+ Requires-Dist: numpy>=1.26
20
+ Requires-Dist: scipy>=1.17.1
21
+ Requires-Dist: xarray>=2026.9.0
22
+ Description-Content-Type: text/markdown
23
+
24
+ # measly🍽️
25
+
26
+ More data helps. measly estimates how much.
27
+
28
+ measly estimates whether a model is data-limited, meaning whether measuring more samples would improve its performance. It is built for fields where each measurement is slow and expensive, such as the life sciences.
29
+
30
+ Note that measly's sweep trains `'n_draws' * 'n_models' * 'n_fractions'` models, which can be extremely computationally expensive. The premise is that this is (should be) still cheaper than collecting new data. If this is not the case, measly isn't appropriate for your use case.
31
+
32
+ ## Problem Definition
33
+
34
+ measly fits a model to many random subsets of the dataset at different fractions of its size. Each draw gives one learning curve, and each curve is fitted with a scaling law. The default is `pow4`:
35
+
36
+ $$L(n) = L_\infty + A\,(n + d)^{-\alpha}$$
37
+
38
+ `pow3`, the three-parameter form most of the literature quotes, is `pow4`
39
+ with $d$ fixed at zero:
40
+
41
+ $$L(n) = L_\infty + A\,n^{-\alpha}$$
42
+
43
+
44
+ ![a learning curve for two model capacities](docs/learning-curve.png)
45
+
46
+ At 75 examples the 75-feature model is more than three times worse. At 750 it
47
+ is 1% better. The curves cross near 400.
48
+
49
+ A pilot study hides that crossover. Measured only at 75 examples, you would
50
+ have picked the small model and been wrong about every later round. A model's
51
+ rank at one sample size says little about its rank at the next.
52
+
53
+ Both gains are small but positive, and the larger model gains about 40% more.
54
+ So collecting more data helps a little, and mostly the high-capacity model.
55
+
56
+ The gain interval is far tighter than the loss interval. Every size
57
+ in a draw is scored on the same test set, so that draw's whole curve shifts up
58
+ or down together. The shift dominates the loss band and cancels in the gain,
59
+ because each curve is compared against its own level, not an average over
60
+ draws.
61
+
62
+ ## Does it work?
63
+
64
+ The example data is synthetic, so the projection can be checked against the
65
+ truth. [`examples/capacity.py`](examples/capacity.py) draws extra rows from the
66
+ same process, refits at four sizes past the measured range, and scores on
67
+ 20,000 held-out rows. These are the open diamonds in the figure. The larger
68
+ sets extend the pilot, as collecting more data would. Only the new rows are
69
+ redrawn across the ten repeats.
70
+
71
+ ```
72
+ model examples actual projected interval error
73
+ ridge(20 features) 938 1.1592 +-0.0027 1.0795 [0.9534, 1.2388] 6.9%
74
+ ridge(20 features) 1125 1.1539 +-0.0047 1.0762 [0.9357, 1.2321] 6.7%
75
+ ridge(20 features) 1312 1.1519 +-0.0047 1.0737 [0.9322, 1.2308] 6.8%
76
+ ridge(20 features) 1500 1.1499 +-0.0042 1.0717 [0.9312, 1.2301] 6.8%
77
+ ridge(75 features) 938 1.0736 +-0.0091 1.1035 [0.9201, 1.3113] 2.8%
78
+ ridge(75 features) 1125 1.0594 +-0.0090 1.1006 [0.9139, 1.3095] 3.9%
79
+ ridge(75 features) 1312 1.0515 +-0.0091 1.0963 [0.9079, 1.3087] 4.3%
80
+ ridge(75 features) 1500 1.0466 +-0.0057 1.0961 [0.9040, 1.3084] 4.7%
81
+ ```
82
+
83
+ All eight land inside their intervals, never more than 6.9% from the median.
84
+
85
+ The projected loss shows no consistent bias. Over eight pilots the projection
86
+ at 2x sat above what actually happened in 3 of 8 for the nearly flat model and
87
+ 5 of 8 for the steeply descending one. The mean error was +0.4% and +2.4%, and
88
+ all sixteen intervals held the truth. The scatter comes from scoring on a slice
89
+ of your own pilot.
90
+
91
+ The audit stops at twice the measured pool. The held-back check fits on the
92
+ smaller sizes and predicts the largest two, a stretch of about 1.9x, so that
93
+ is as far as the method carries evidence. Projecting ten times out would
94
+ return a number and support none of it.
95
+
96
+ The error bars on the diamonds are smaller than the markers. Ten independent
97
+ collections at one size vary by 0.0027 to 0.0091, while the band spans about
98
+ 0.3, so the band is 30 to 100 times wider. That is the safe direction to err
99
+ in. Most of the width is test-set noise, not uncertainty about the curve.
100
+
101
+ Subsets are drawn without replacement. Each draw resplits the pool, so even the
102
+ full-size fit varies across draws. A bootstrap would add little to that spread
103
+ and would bias the level up, because a resample repeats rows.
104
+
105
+ `analyse` always checks itself. It refits the curves on the smaller sample
106
+ sizes, predicts the largest ones you already measured, and reports the error.
107
+ Above 15% the projection MUST NOT be relied on: the shape fits your data but
108
+ does not predict it.
109
+
110
+ `analyse(...).results` is an `xarray.DataArray` over `(fraction, model, draw)`
111
+ if you want the raw measurements.
112
+
113
+ ## Previous work in this area/Inspiration
114
+
115
+ - [The Shape of Learning Curves: a Review](https://arxiv.org/abs/2103.10948). Viering & Loog, 2021.
116
+ - [Deep Learning Scaling is Predictable, Empirically](https://arxiv.org/abs/1712.00409). Hestness et al., 2017.
117
+ - [Scaling Laws for Neural Language Models](https://arxiv.org/abs/2001.08361). Kaplan et al., 2020.
118
+ - [Training Compute-Optimal Large Language Models](https://arxiv.org/abs/2203.15556). Hoffmann et al., 2022.
119
+ - [Revisiting Neural Scaling Laws in Language and Vision](https://arxiv.org/abs/2209.06640). Alabdulmohsin et al., 2022.
120
+ - [Broken Neural Scaling Laws](https://github.com/ethancaballero/broken_neural_scaling_laws). Caballero et al., 2022.
121
+ - [How Much More Data Do I Need?](https://research.nvidia.com/labs/toronto-ai/estimatingrequirements/). Mahmood et al., 2022.
122
+ - [Estimation of Predictive Performance in High-Dimensional Data Settings using Learning Curves](https://arxiv.org/abs/2206.03825). Goedhart et al., 2022.
@@ -0,0 +1,10 @@
1
+ measly/__init__.py,sha256=YqjYyhpvPHzh_IQHABbyxb-MGSiKAAiMOXSiFVuHgFk,824
2
+ measly/analysis.py,sha256=vA2gtltqK4Wuz9PCF7inMfruClgFZWkMGSV9kycv5EI,10327
3
+ measly/fit.py,sha256=SZ362l8Td5Bw_qj-2VFcGmcpy_xVcdZfOePwFFNpN_c,4722
4
+ measly/interfaces.py,sha256=S65JCnXymA8mzoHPxQZrNwquiC3hq1WnoZiVIu-XwBI,1143
5
+ measly/plot.py,sha256=8-JZKUEecyQhjpqj557TuyiVVre2NOVe2iHnZZyvO_c,3670
6
+ measly/sweep.py,sha256=1UjTRsZrzMFqIEPfr94vVCqa6xOuTc0IKl9Y2NnpCv4,3726
7
+ measly-0.1.0.dist-info/METADATA,sha256=Oz2yDTJ-Y0BFJLhPoPgR2KUKUmNGZPyQewDXFnc9NfI,6521
8
+ measly-0.1.0.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
9
+ measly-0.1.0.dist-info/licenses/LICENSE,sha256=IQO1FLCH5dypLlWDpWMdoQe4U2xcERhaZAOeeF6fKGE,1068
10
+ measly-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.4
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Remi Tregan
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.