eosframes 1.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.
eosframes/scale.py ADDED
@@ -0,0 +1,1713 @@
1
+ """Type-aware robust scaler for Ersilia model outputs.
2
+
3
+ Each numeric feature column is auto-classified at fit time as one of
4
+ ``constant``, ``binary``, ``count``, or ``continuous`` (with a further
5
+ ``right`` / ``left`` / ``centered`` sub-case for continuous and count
6
+ columns that have a non-zero mode), and a type-specific transform is
7
+ recorded. Output values land in a per-type region whose maximum spread
8
+ is 1, so columns are commensurable for distance calculations:
9
+
10
+ * ``constant`` → output 0.
11
+ * ``binary`` → ``{0, 1}`` (lower → 0, higher → 1).
12
+ * ``count`` (mode 0) → ``[0, 1]``, with 0 pinned at 0.
13
+ Sparse mode-0 columns whose output collapses to a few distinct
14
+ values get a ``degenerate: True`` advisory flag (warn-only).
15
+ * ``count`` (mode ≠ 0) → per-side linear body + ``tanh``
16
+ asymptote, same shape as ``continuous_centered``. Anchors:
17
+ ``upper = min(max(Q3 + 1.5·IQR, p99), arr.max())``,
18
+ ``lower = max(min(Q1 − 1.5·IQR, p1), arr.min())``. The
19
+ ``max(Tukey, p99) / min(Tukey, p1)`` part keeps the original
20
+ outlier-robust anchor (whichever side is wider), and the
21
+ ``arr.max() / arr.min()`` cap prevents over-extension past the
22
+ data range. Mode → 0; the linear body maps ``[mode, high_anchor]``
23
+ → ``[0, body_target]`` (and the mirror on the lower side) with
24
+ ``body_target = _CENTERED_BODY_TARGET = 0.5`` — so the bulk of
25
+ count data sits inside ``[-0.5, 0.5]`` and counts past the anchors
26
+ asymptote toward ``±1`` via the slope-continuous ``tanh``. **Soft
27
+ cap on per-side extent ratio**: when
28
+ ``max(upper_span, lower_span) / min(upper_span, lower_span) >
29
+ _COUNT_SHIFTED_EXTENT_RATIO_MAX`` (currently ``2.0``), the wider
30
+ side's extent is capped at ``2 × the narrower`` and the affected
31
+ anchor is recomputed — so heavily right- or left-tailed count
32
+ distributions don't squash their distinct discrete values into the
33
+ body. The excess raw range flows into the ``tanh`` tail past the
34
+ anchor, where each distinct count gets a distinct output in
35
+ ``(±0.5, ±1)``. ``fit_notes`` records ``extent_ratio`` and
36
+ ``extent_capped``. Distinct outliers always get distinct outputs
37
+ (no flat clip plateau); distinct integer counts inside the body
38
+ keep distinct outputs too (no compression there).
39
+ * ``continuous`` right-skewed → triggered when **both**
40
+ ``bowley > 0.3`` *and* the tail-asymmetry ratio
41
+ ``(p99 − median) / (median − p1) > 3``. Bowley alone is a bulk-only
42
+ asymmetry measure and trips on borderline columns that have real
43
+ spread on both sides. Requiring the tails to agree with the
44
+ bulk keeps those columns in the centered branch where their
45
+ bell-ish shape is preserved. The fit itself is a linear body
46
+ ``[arr.min, body_anchor] → [0, body_target]`` plus a smooth tail:
47
+
48
+ * **Heavy tail** (``p99 > Tukey``): 3-segment piecewise (``tail =
49
+ "piecewise"``). Linear body ``[arr.min, Tukey] → [0, body_target
50
+ = 0.8]`` keeps the bulk visible (Q3 lands at ``≈ 0.35``); a
51
+ slower linear "middle" segment ``[Tukey, p99] → [body_target,
52
+ mid_target]`` spreads the upper-decile outliers; a finite
53
+ ``_quadratic_tail`` reaches ``(arr.max, 1)`` with zero slope at
54
+ ``arr.max``. ``mid_target`` is derived for ``C¹`` smoothness at
55
+ the middle→tail join:
56
+ ``mid_target = (2·a + b·body_target) / (2·a + b)``
57
+ with ``a = p99 − Tukey`` and ``b = arr.max − p99``. The
58
+ histogram shows the bulk in ``[0, body_target]`` and a separate
59
+ smaller "tail bump" in ``[body_target, 1]`` — no spike at ``+1``.
60
+ * **Light tail** (``p99 ≤ Tukey``): finite quadratic from
61
+ ``(body_anchor, body_target)`` to ``(arr.max, 1)`` with zero slope
62
+ at ``arr.max`` and slope continuity at the body boundary.
63
+ ``body_anchor`` is capped at ``arr.max`` so bounded distributions
64
+ fill the full output range, and
65
+ ``body_target`` is derived from the geometry so the quadratic
66
+ reaches exactly ``1`` at ``arr.max``. Only OOD inputs past
67
+ ``arr.max`` clip at ``1``.
68
+ * ``continuous`` left-skewed → mirror of right.
69
+ * ``continuous`` centered → per-side linear body +
70
+ finite-reach quadratic tail in ``[-1, 1]``. The per-side body
71
+ extents and max extents adapt independently:
72
+
73
+ ``upper_body_extent = max(scale, min(p99 − center, tukey_upper − center))``
74
+ ``lower_body_extent = max(scale, min(center − p1, center − tukey_lower))``
75
+ ``upper_max_extent = max(upper_body_extent, arr.max − center)``
76
+ ``lower_max_extent = max(lower_body_extent, center − arr.min)``
77
+
78
+ Tukey caps ``p99`` / ``p1`` to stay outlier-robust;
79
+ ``scale = half-IQR`` floors the body so tight clean data still has a
80
+ usable body. ``body_target`` is derived per side at apply time by
81
+ :func:`_light_tail_body_target` and then **capped at**
82
+ :data:`_CENTERED_BODY_TARGET` (``0.5``), so the bulk of the density
83
+ visually lives in ``[-0.5, 0.5]`` to match ``count_shifted`` and stay
84
+ commensurable with the one-sided skewed kinds. **Dually**, the
85
+ effective ``max_extent`` used for the tail is floored at
86
+ :data:`_CENTERED_MAX_EXTENT_RATIO` × ``body_extent`` (``3·b``); the
87
+ two constants are locked together because the body→tail join is
88
+ ``C¹``-smooth precisely when ``max_extent = 3·body_extent`` given
89
+ ``body_target = 0.5`` (body slope ``0.5/b`` equals quadratic-tail
90
+ slope ``1/(a−b)``). The cap activates when ``max_extent <
91
+ 3·body_extent`` (typical Gaussian-ish or bimodal data with little
92
+ tail past the body); past that threshold the natural derivation
93
+ produces ``body_target < 0.5`` already and both the cap and the
94
+ floor are no-ops. Past the body, :func:`_quadratic_tail` reaches
95
+ ``±1`` exactly at ``±effective_max`` (= ``±3·body_extent`` when
96
+ the floor is active, otherwise the actual data edge on that side);
97
+ inputs past that clip to ``±1``. Distinct outlier values spread
98
+ monotonically across ``[body_target, 1]`` — no visual pile-up at
99
+ ``±1`` and no slope kink at the body→tail join.
100
+
101
+ **Safety gate.** The per-side design is only safe when both sides
102
+ are similarly populated. When ``max(upper, lower) / min(upper,
103
+ lower) > _CENTERED_BODY_RATIO_MAX`` (currently 2.0) the
104
+ distribution is effectively one-sided against a hard bound (a
105
+ probability output piled at 0, say). The tight side would
106
+ otherwise get a steep body slope that distorts the histogram. The
107
+ fit collapses to a **symmetric** body — both sides equal to
108
+ ``max(2·scale, min(upper, lower))``, the IQR-floored smaller side
109
+ — AND a **symmetric** ``max_extent`` (the larger of the two sides)
110
+ so both sides share the same derived ``body_target`` and the
111
+ centre has no density step. ``fit_notes`` records ``body_ratio``
112
+ and ``symmetric_fallback`` so the gate can be audited after the
113
+ fact.
114
+
115
+ Output dtype is a **transform-time choice**, not a fit-time one. The
116
+ scaler JSON contains only dtype-agnostic parameters; pass
117
+ ``output_dtype`` to :func:`transform` / :func:`transform_file` (or use
118
+ ``--quantize`` on the CLI) to select between ``float32`` (default,
119
+ NaN-preserving) and ``int8``. The int8 mapping is per-column: every
120
+ column's documented output region maps linearly to ``[-127, 127]`` so
121
+ asymmetric regions (binary, count, right, left) use the full int8
122
+ range, not just half of it. Sentinel ``-128`` is reserved for NaN.
123
+ :func:`fit_file` accepts ``output_dtype`` only as a convenience for
124
+ the inline fit-then-transform path it offers via its ``output_path``
125
+ argument.
126
+
127
+ Every numeric feature column is fitted — there is no missing-value
128
+ skip threshold. Columns that have zero non-NaN values fit as
129
+ ``type: constant, value: 0.0`` (every row transforms to 0 for non-NaN
130
+ inputs and NaN for NaN inputs).
131
+
132
+ Each per-column params dict also carries an ``impute_value`` recorded
133
+ at fit time: the median of the column's non-NaN training values,
134
+ rounded to ``int`` if the column was integer-valued and kept as
135
+ ``float`` otherwise (all-NaN columns get ``0.0``). The transform
136
+ ignores it by default — NaN propagates through to the output, just as
137
+ before. Passing ``impute=True`` (or ``--impute`` on the CLI) replaces
138
+ every input NaN with the recorded ``impute_value`` before dispatch,
139
+ so the output column has no NaN (and, under ``--quantize``, no
140
+ ``-128`` sentinels).
141
+
142
+ Schema
143
+ ------
144
+
145
+ A fitted scaler is written as a flat JSON envelope plus a ``columns``
146
+ dict keyed by feature name. Each column entry has at most three keys:
147
+
148
+ * ``transform`` — the transform-time payload. Always carries a
149
+ ``kind`` string discriminator; the rest of the keys depend on the
150
+ kind (see ``_OUTPUT_REGIONS`` for the seven kinds and their
151
+ per-kind documented output region).
152
+ * ``impute_value`` — the median-with-dtype-rounding fill for
153
+ ``transform(..., impute=True)``. Peer of ``transform`` so the
154
+ imputation logic stays separable from the dispatch math.
155
+ * ``fit_notes`` — *optional*. Diagnostics never read at transform
156
+ time (e.g. ``bowley`` for continuous, ``degenerate`` /
157
+ ``mode_fraction`` for sparse zero-mode counts, ``scale`` /
158
+ ``scale_kind`` for centered continuous). Omitted entirely when
159
+ empty.
160
+
161
+ Output region is **derived** at quantize time from the ``kind``, not
162
+ stored — drift-free.
163
+
164
+ The envelope also records ``eosframes_version`` (the full running
165
+ ``eosframes.__version__`` that produced the JSON), plus ``method``,
166
+ ``model_id``, ``model_version``, ``fitted_at``, and ``n_rows``.
167
+ :func:`transform_file` rejects any scaler whose recorded
168
+ ``eosframes_version`` doesn't share the **major** component with the
169
+ running version (so the whole ``0.x.y`` line is mutually compatible
170
+ and only a bump to ``1.0.0`` invalidates older scalers). There is no
171
+ hand-maintained schema number.
172
+ """
173
+
174
+ import json
175
+ import os
176
+ from datetime import datetime
177
+ from typing import Callable, Dict, Optional, Tuple
178
+
179
+ import h5py
180
+ import numpy as np
181
+ import pandas as pd
182
+
183
+ from . import __version__ as _PACKAGE_VERSION
184
+ from .exceptions import EosframesError
185
+ from .logger import get_logger
186
+ from .naming import (
187
+ is_valid_name,
188
+ is_valid_transformer_name,
189
+ parse_name,
190
+ parse_transformer_name,
191
+ )
192
+
193
+ _META_COLS = {"key", "input"}
194
+
195
+ _METHOD_NAME = "robust_typed"
196
+
197
+ _VALID_OUTPUT_DTYPES = ("float32", "int8")
198
+ _DEFAULT_OUTPUT_DTYPE = "float32"
199
+
200
+ _INT8_NAN_SENTINEL = -128
201
+ _INT8_MAX_VAL = 127
202
+
203
+ _BOWLEY_THRESHOLD = 0.3
204
+
205
+ # A column counts as right- or left-skewed only when *both* the bulk
206
+ # (Bowley) and the tails (p1, p99 relative to median) agree on the
207
+ # direction. Bowley alone is a bulk-only asymmetry measure and trips
208
+ # on borderline columns whose tails happen to have comparable extent
209
+ # on both sides; requiring the tail-asymmetry ratio to exceed this
210
+ # threshold sends those into the centered branch where their
211
+ # bell-ish shape is preserved.
212
+ _TAIL_ASYMMETRY_THRESHOLD = 3.0
213
+ _TUKEY_THRESHOLD = 3.0
214
+ _LINEAR_CLIP_DIVISOR_MULTIPLIER = 2.0
215
+
216
+ # F4: a mode-0 count is flagged degenerate when it is mostly the mode and
217
+ # its scaled output collapses to a handful of distinct values. The flag is
218
+ # advisory — transform behavior is unchanged.
219
+ _DEGENERATE_DISTINCT_MAX = 4
220
+ _DEGENERATE_MODE_FRACTION_MIN = 0.4
221
+
222
+ # Mode-0 count columns anchor at the larger of the Tukey whisker
223
+ # (`q3 + 1.5·IQR`) and the 99th percentile of the data. On heavy-zero
224
+ # distributions Tukey sits inside the realistic tail (Q3 and IQR collapse
225
+ # toward 0), so the p99 lifts the anchor; on denser counts Tukey is the
226
+ # more generous of the two and it wins. Both are robust to a single
227
+ # rogue outlier (unlike `arr.max()`).
228
+ _COUNT_HIGH_PERCENTILE = 0.99
229
+
230
+ # Continuous fits use a linear body + smooth tail mapping. The body is
231
+ # a single linear segment that ends at ``±body_target``; past the body,
232
+ # the tail is either an asymptotic tanh (heavy tail — data extends past
233
+ # the body anchor) or a finite quadratic ending at ``±1`` with zero
234
+ # slope (light tail — data has a real ceiling near the body anchor).
235
+ # Both tails preserve slope continuity at the body boundary, so the
236
+ # overall transform has no kinks and distinct inputs in the tail get
237
+ # distinct outputs — no collisions at a clip plateau.
238
+ _CONTINUOUS_BODY_TARGET = 0.8
239
+
240
+ # Two-sided body target. ``continuous_centered`` and ``count_shifted``
241
+ # both produce output in ``[-1, 1]`` symmetrically around 0, so the
242
+ # natural visual split is body 50% / each tail 25%. A reader sees the
243
+ # bulk inside ``[-0.5, 0.5]`` and anything past ``±0.5`` reads
244
+ # unambiguously as tail / outlier. ``count_shifted`` uses this as a
245
+ # **hard target** (every fit lands the body at ``±0.5``); for
246
+ # ``continuous_centered`` it is a **cap** on the per-side geometry-
247
+ # derived target, so heavy-tailed distributions can still pick a
248
+ # smaller body_target while light-tailed and bimodal fits collapse to
249
+ # the shared 50%-body convention rather than sprawling across
250
+ # ``[-1, 1]``. The ``tanh`` (count_shifted) or quadratic (centered)
251
+ # tail past the body stretches outliers across the remaining half, so
252
+ # int8-quantized outliers get plenty of distinct levels and centered
253
+ # columns stay commensurable with the one-sided skewed kinds for
254
+ # distance metrics.
255
+ _CENTERED_BODY_TARGET = 0.5
256
+
257
+ # Minimum ``max_extent / body_extent`` ratio for ``continuous_centered``.
258
+ # Paired with :data:`_CENTERED_BODY_TARGET` = 0.5: the body→tail join
259
+ # is ``C¹``-smooth precisely when ``max_extent = 3·body_extent`` (body
260
+ # slope ``0.5/b`` equals quadratic-tail slope ``1/(a−b)``). When the
261
+ # data provides a smaller ratio — bounded distributions like a U-shape
262
+ # beta, truncated normal, narrow bimodal, or any column where the
263
+ # Tukey reach already touches the data edge — the effective
264
+ # ``max_extent`` is floored at ``3·body_extent`` at apply time so the
265
+ # quadratic tail has enough raw domain to be gentle. Without the
266
+ # floor, the few real points sitting in the tiny ``[body_extent,
267
+ # max_extent]`` raw band would get sprayed across ``(0.5, 1.0]``;
268
+ # with it, they collapse to a thin pile-up just past ``±0.5`` and OOD
269
+ # inputs still reach ``±1`` only past ``±3·body_extent`` from the
270
+ # centre. The constant is **mathematically locked** to the cap value:
271
+ # ``_CENTERED_MAX_EXTENT_RATIO = 2/_CENTERED_BODY_TARGET − 1``. If
272
+ # the cap ever moves, this constant must move with it.
273
+ _CENTERED_MAX_EXTENT_RATIO = 3.0
274
+
275
+ # Inside ``continuous_centered``: when the per-side body extents differ
276
+ # by more than this factor, the fit falls back to a symmetric body
277
+ # (both sides equal to the smaller extent). Per-side asymmetry past
278
+ # this point comes from a one-sided bounded distribution (a probability
279
+ # output piled near 0, for instance), where the tight side ends up with
280
+ # a steep, distorting body slope at the centre. The cap of 2.0 keeps
281
+ # genuinely asymmetric bells per-side while folding pile-up columns
282
+ # back to a symmetric fit.
283
+ _CENTERED_BODY_RATIO_MAX = 2.0
284
+
285
+ # Max allowed ratio between the per-side body extents in
286
+ # ``count_shifted`` (mode-nonzero count fit). When one side's extent
287
+ # exceeds ``_COUNT_SHIFTED_EXTENT_RATIO_MAX × the other``, the wider
288
+ # side is **capped** at that ratio (a soft cap — the lower side stays
289
+ # unchanged in the typical right-tailed case). The excess raw range
290
+ # on the wider side then flows into the existing slope-continuous
291
+ # ``tanh`` tail past the anchor, so distinct outlier counts get
292
+ # distinct outputs in ``(0.5, ~1)`` instead of being squashed into
293
+ # the body's ``[0, 0.5]``. The value mirrors
294
+ # :data:`_CENTERED_BODY_RATIO_MAX` for design parity between the two
295
+ # centered-style kinds; this is a softer enforcement than centered's
296
+ # symmetric-collapse fallback because ``count_shifted`` already has
297
+ # an asymptotic tail to absorb the excess.
298
+ _COUNT_SHIFTED_EXTENT_RATIO_MAX = 2.0
299
+
300
+ # Half-width of the Hermite smoothstep blend window at the body→middle
301
+ # junction of the 3-segment piecewise right/left-skew fit, as a fraction
302
+ # of ``min(body_span, mid_span)``. The slope on either side of the
303
+ # junction differs by ~20× for heavy-tailed columns; without smoothing,
304
+ # the density jump at the junction shows up as a visible "second hill"
305
+ # in the scaled histogram. The blend window covers half of each
306
+ # neighbouring segment by default — wide enough to spread the slope
307
+ # change across most of the segment-2 row mass, so the visible bump
308
+ # softens into a continuous shoulder.
309
+ _PIECEWISE_BLEND_FRACTION = 0.5
310
+
311
+
312
+ def _tanh_tail(
313
+ mag: np.ndarray, body_extent: float, body_target: float
314
+ ) -> np.ndarray:
315
+ """Asymptotic tail for ``|x - center| > body_extent``.
316
+
317
+ Returns ``y = body_target + (1 − body_target)·tanh(c·u)`` for
318
+ ``u = mag − body_extent ≥ 0``, where
319
+ ``c = body_target / (body_extent · (1 − body_target))`` makes the
320
+ tail slope at ``u = 0`` exactly equal to the body slope
321
+ ``body_target / body_extent`` — so the body→tail join is
322
+ ``C¹``-smooth, no kink, no smoothstep blend needed. As ``u → ∞``,
323
+ ``y → 1`` asymptotically; distinct inputs always get distinct
324
+ outputs.
325
+ """
326
+ if body_target >= 1.0 or body_extent <= 0:
327
+ return np.full_like(mag, body_target, dtype=float)
328
+ c = body_target / (body_extent * (1.0 - body_target))
329
+ return body_target + (1.0 - body_target) * np.tanh(c * (mag - body_extent))
330
+
331
+
332
+ def _quadratic_tail(
333
+ mag: np.ndarray, body_extent: float, body_target: float, max_extent: float
334
+ ) -> np.ndarray:
335
+ """Finite-reach tail for ``|x - center|`` in ``[body_extent, max_extent]``.
336
+
337
+ Reaches ``y = 1`` exactly at ``mag = max_extent`` with zero slope, and
338
+ matches ``(body_extent, body_target)`` with body slope at the body
339
+ boundary. The slope-continuity + zero-slope-at-end constraints fix
340
+ a unique quadratic; the caller is responsible for picking
341
+ ``body_target`` so the quadratic is monotone (in practice
342
+ ``body_target = 2·body_extent / (body_extent + max_extent)``).
343
+
344
+ Values past ``max_extent`` clip to ``1``.
345
+ """
346
+ if body_extent <= 0 or max_extent <= body_extent:
347
+ return np.full_like(mag, 1.0, dtype=float)
348
+ alpha = (1.0 - body_target) / (max_extent - body_extent) ** 2
349
+ y = 1.0 - alpha * (max_extent - mag) ** 2
350
+ return np.clip(y, body_target, 1.0)
351
+
352
+
353
+ def _light_tail_body_target(body_extent: float, max_extent: float) -> float:
354
+ """Derive ``body_target`` for a finite (quadratic) tail.
355
+
356
+ Picked so the body→tail join is ``C¹``-smooth: the linear body
357
+ slope ``body_target / body_extent`` equals the quadratic's slope at
358
+ its left edge. Solving gives ``body_target = 2·b / (b + a)`` where
359
+ ``b = body_extent`` and ``a = max_extent``. ``b = a`` (no tail)
360
+ yields ``body_target = 1`` and the quadratic degenerates to a
361
+ constant ``1`` past the body — same as a hard clip in that
362
+ degenerate case.
363
+ """
364
+ if max_extent <= 0:
365
+ return 1.0
366
+ return 2.0 * body_extent / (body_extent + max_extent)
367
+
368
+
369
+ # Per-kind documented output region. Read by :func:`_quantize_to_int8`
370
+ # to map each column's region linearly into the int8 range
371
+ # ``[-127, 127]`` (sentinel ``-128`` reserved for NaN). Single source of
372
+ # truth — never stored in the scaler JSON, derived at runtime from
373
+ # ``entry["transform"]["kind"]``.
374
+ _OUTPUT_REGIONS: Dict[str, Tuple[float, float]] = {
375
+ "constant": (0.0, 0.0),
376
+ "binary": (0.0, 1.0),
377
+ "count_zero_mode": (0.0, 1.0),
378
+ "count_shifted": (-1.0, 1.0),
379
+ "continuous_right_skew": (0.0, 1.0),
380
+ "continuous_left_skew": (-1.0, 0.0),
381
+ "continuous_centered": (-1.0, 1.0),
382
+ }
383
+
384
+
385
+ def _major(v: str) -> str:
386
+ """Return the major component of a PEP 440-style version string.
387
+
388
+ ``"0.1.0"`` → ``"0"``, ``"1.2.3a1"`` → ``"1"``, ``"unknown"`` →
389
+ ``"unknown"`` (anything without a ``.`` is returned verbatim).
390
+ Used by :func:`transform_file` so scalers fitted by any
391
+ ``eosframes`` release in the same major line stay valid.
392
+ """
393
+ return v.split(".", 1)[0] if v else v
394
+
395
+
396
+ # ---------------------------------------------------------------------------
397
+ # Type classification and per-type fit
398
+ # ---------------------------------------------------------------------------
399
+
400
+
401
+ def _is_integer_valued(series: pd.Series) -> bool:
402
+ arr = series.dropna().to_numpy(dtype=float)
403
+ if arr.size == 0:
404
+ return False
405
+ return bool(np.all(np.isfinite(arr)) and np.allclose(arr, np.round(arr)))
406
+
407
+
408
+ def _mode_value(series: pd.Series) -> float:
409
+ """Most-common non-NaN value, ties broken by the lowest value."""
410
+ counts = series.dropna().value_counts(sort=False)
411
+ if counts.empty:
412
+ raise EosframesError("Cannot compute mode of an all-NaN column.")
413
+ max_count = counts.max()
414
+ candidates = counts[counts == max_count].index.tolist()
415
+ return float(min(candidates))
416
+
417
+
418
+ def _classify_type(series: pd.Series) -> str:
419
+ non_nan = series.dropna()
420
+ n_unique = non_nan.nunique()
421
+ if n_unique <= 1:
422
+ return "constant"
423
+ if n_unique == 2:
424
+ return "binary"
425
+ if _is_integer_valued(series) and float(non_nan.min()) >= 0.0:
426
+ return "count"
427
+ return "continuous"
428
+
429
+
430
+ def _compute_impute_value(series: pd.Series):
431
+ """Median of *series* with dtype-respecting rounding.
432
+
433
+ Integer-valued training columns return an ``int`` (median rounded
434
+ to the nearest integer) so a later ``impute=True`` transform never
435
+ writes fractional values into an originally-int column. Float
436
+ columns return a ``float``. All-NaN columns return ``0.0``.
437
+ """
438
+ non_nan = series.dropna()
439
+ if non_nan.empty:
440
+ return 0.0
441
+ median = float(non_nan.median())
442
+ if _is_integer_valued(series):
443
+ return int(round(median))
444
+ return median
445
+
446
+
447
+ def _compute_robust_scale(
448
+ series: pd.Series,
449
+ ) -> Tuple[float, float, str]:
450
+ """Median + scale via half-IQR / MAD / range cascade.
451
+
452
+ Returns ``(center, scale, scale_kind)``. ``scale_kind`` is one of
453
+ ``"half_iqr"``, ``"mad"``, ``"range"``. Raises ``EosframesError`` if
454
+ the column is effectively constant (all three scales are 0).
455
+ """
456
+ arr = series.dropna().to_numpy(dtype=float)
457
+ median = float(np.median(arr))
458
+ q1, q3 = np.quantile(arr, [0.25, 0.75])
459
+ half_iqr = float((q3 - q1) / 2.0)
460
+ if half_iqr > 0:
461
+ return median, half_iqr, "half_iqr"
462
+ mad = float(np.median(np.abs(arr - median)))
463
+ mad_scale = mad * 1.4826
464
+ if mad_scale > 0:
465
+ return median, mad_scale, "mad"
466
+ half_range = float((arr.max() - arr.min()) / 2.0)
467
+ if half_range > 0:
468
+ return median, half_range, "range"
469
+ raise EosframesError("Column is effectively constant; route to constant branch.")
470
+
471
+
472
+ def _compute_bowley(series: pd.Series, scale_kind: str) -> float:
473
+ """Bowley skewness ``((Q3 + Q1) - 2·median) / IQR``.
474
+
475
+ Returns 0.0 when the half-IQR cascade did not apply (skew is
476
+ ill-defined for those degenerate cases).
477
+ """
478
+ if scale_kind != "half_iqr":
479
+ return 0.0
480
+ arr = series.dropna().to_numpy(dtype=float)
481
+ q1, q2, q3 = np.quantile(arr, [0.25, 0.5, 0.75])
482
+ iqr = q3 - q1
483
+ if iqr <= 0:
484
+ return 0.0
485
+ return float(((q3 + q1) - 2.0 * q2) / iqr)
486
+
487
+
488
+ def _fit_constant(series: pd.Series) -> dict:
489
+ # ``constant`` columns always emit 0 for non-NaN at transform time;
490
+ # the original training value lives in ``impute_value`` if anyone
491
+ # needs to recover it. The transform payload itself is empty.
492
+ return {
493
+ "transform": {"kind": "constant"},
494
+ "impute_value": _compute_impute_value(series),
495
+ }
496
+
497
+
498
+ def _fit_binary(series: pd.Series) -> dict:
499
+ uniques = sorted(float(v) for v in series.dropna().unique())
500
+ return {
501
+ "transform": {"kind": "binary", "low": uniques[0], "high": uniques[1]},
502
+ "impute_value": _compute_impute_value(series),
503
+ }
504
+
505
+
506
+ def _fit_continuous_centered(series: pd.Series) -> dict:
507
+ """Fit a centered continuous column with per-side body + quadratic tail.
508
+
509
+ Linear body to ``±body_target`` on each side, then a finite-reach
510
+ ``_quadratic_tail`` to ``±1`` at each side's actual data edge — so
511
+ outliers spread monotonically across ``[body_target, 1]`` instead
512
+ of saturating asymptotically and visually piling at ``±1``.
513
+
514
+ * Per-side body extent: ``upper_body_extent =
515
+ max(scale, min(p99 − center, tukey_upper − center))`` and the
516
+ mirror for the lower side. Tukey caps p99 / p1 when outliers
517
+ infiltrate them; ``scale = half-IQR`` floors the body so tight
518
+ clean data still has a usable body.
519
+ * Per-side ``max_extent`` is the actual data extreme on that side
520
+ (``arr.max − center`` for upper, ``center − arr.min`` for lower),
521
+ floored at ``body_extent`` so the quadratic tail is always
522
+ well-defined.
523
+ * ``body_target`` is derived per side at apply time by
524
+ :func:`_light_tail_body_target` so the body→tail join is
525
+ ``C¹``-smooth. It's not stored — derivable from
526
+ ``(body_extent, max_extent)`` alone.
527
+ * **Safety gate.** When the body-extent ratio ``max / min``
528
+ exceeds ``_CENTERED_BODY_RATIO_MAX``, the distribution is
529
+ effectively one-sided against a hard bound (e.g. ``clintox``).
530
+ The fit collapses to a **symmetric** body (smaller side, floored
531
+ at the IQR) AND a **symmetric** ``max_extent`` (max of both
532
+ sides), so both sides share the same derived ``body_target`` and
533
+ there's no density step at the centre.
534
+ """
535
+ center, scale, scale_kind = _compute_robust_scale(series)
536
+ arr = series.dropna().to_numpy(dtype=float)
537
+ q1, q3 = np.quantile(arr, [0.25, 0.75])
538
+ iqr = float(q3 - q1)
539
+ tukey_upper = float(q3 + 1.5 * iqr)
540
+ tukey_lower = float(q1 - 1.5 * iqr)
541
+ p_upper = float(np.quantile(arr, _COUNT_HIGH_PERCENTILE))
542
+ p_lower = float(np.quantile(arr, 1.0 - _COUNT_HIGH_PERCENTILE))
543
+
544
+ upper_body_extent = max(scale, min(p_upper - center, tukey_upper - center))
545
+ lower_body_extent = max(scale, min(center - p_lower, center - tukey_lower))
546
+
547
+ upper_max_extent = max(upper_body_extent, float(arr.max()) - center)
548
+ lower_max_extent = max(lower_body_extent, center - float(arr.min()))
549
+
550
+ denom = max(min(upper_body_extent, lower_body_extent), 1e-12)
551
+ body_ratio = max(upper_body_extent, lower_body_extent) / denom
552
+ symmetric_fallback = body_ratio > _CENTERED_BODY_RATIO_MAX
553
+ if symmetric_fallback:
554
+ # Symmetric body (smaller side, IQR-floored) and symmetric
555
+ # max_extent (larger side's reach) — same derived body_target
556
+ # on both sides, no centre density step.
557
+ sym_body = max(2.0 * scale, min(upper_body_extent, lower_body_extent))
558
+ sym_max = max(upper_max_extent, lower_max_extent, sym_body)
559
+ upper_body_extent = lower_body_extent = sym_body
560
+ upper_max_extent = lower_max_extent = sym_max
561
+
562
+ return {
563
+ "transform": {
564
+ "kind": "continuous_centered",
565
+ "center": float(center),
566
+ "upper_body_extent": float(upper_body_extent),
567
+ "lower_body_extent": float(lower_body_extent),
568
+ "upper_max_extent": float(upper_max_extent),
569
+ "lower_max_extent": float(lower_max_extent),
570
+ },
571
+ "impute_value": _compute_impute_value(series),
572
+ "fit_notes": {
573
+ "scale": float(scale),
574
+ "scale_kind": scale_kind,
575
+ "body_ratio": float(body_ratio),
576
+ "symmetric_fallback": bool(symmetric_fallback),
577
+ },
578
+ }
579
+
580
+
581
+ def _fit_continuous(series: pd.Series) -> dict:
582
+ center, scale, scale_kind = _compute_robust_scale(series)
583
+ bowley = _compute_bowley(series, scale_kind)
584
+ arr = series.dropna().to_numpy(dtype=float)
585
+ q1, q3 = np.quantile(arr, [0.25, 0.75])
586
+ iqr = float(q3 - q1)
587
+ # Tail asymmetry: right side vs left side relative to the median.
588
+ # Both Bowley (bulk) and tail asymmetry must agree on the direction
589
+ # to commit to the right/left-skew branch — otherwise centered.
590
+ median_val = float(np.quantile(arr, 0.5))
591
+ p1_val = float(np.quantile(arr, 1.0 - _COUNT_HIGH_PERCENTILE))
592
+ p99_val = float(np.quantile(arr, _COUNT_HIGH_PERCENTILE))
593
+ _eps = 1e-12
594
+ right_span = max(p99_val - median_val, 0.0)
595
+ left_span = max(median_val - p1_val, 0.0)
596
+ tail_asymmetry_right = right_span / max(left_span, _eps)
597
+ tail_asymmetry_left = left_span / max(right_span, _eps)
598
+ is_right_skew = (
599
+ bowley > _BOWLEY_THRESHOLD
600
+ and tail_asymmetry_right > _TAIL_ASYMMETRY_THRESHOLD
601
+ )
602
+ is_left_skew = (
603
+ bowley < -_BOWLEY_THRESHOLD
604
+ and tail_asymmetry_left > _TAIL_ASYMMETRY_THRESHOLD
605
+ )
606
+
607
+ if is_right_skew:
608
+ tukey_upper = float(q3 + 1.5 * iqr)
609
+ p99_upper = float(np.quantile(arr, _COUNT_HIGH_PERCENTILE))
610
+ low = float(arr.min())
611
+ arr_max = float(arr.max())
612
+ if p99_upper > tukey_upper:
613
+ # Heavy tail: 3-segment piecewise — linear body to
614
+ # (body_anchor=Tukey, body_target=0.8), linear "middle" to
615
+ # (mid_anchor=p99, mid_target), then quadratic tail to
616
+ # (high_anchor=arr.max, 1). mid_target is derived from a
617
+ # C¹-smoothness constraint at the middle→tail join so the
618
+ # quadratic stays monotone. The bulk lives in segment 1
619
+ # (visible), the outliers spread linearly across segments
620
+ # 2-3 instead of saturating at +1.
621
+ body_anchor = float(tukey_upper)
622
+ mid_anchor = float(p99_upper)
623
+ high_anchor = float(arr_max)
624
+ body_target = float(_CONTINUOUS_BODY_TARGET)
625
+ a = mid_anchor - body_anchor
626
+ b = high_anchor - mid_anchor
627
+ if b <= 0.0:
628
+ mid_target = 1.0
629
+ else:
630
+ mid_target = (2.0 * a + b * body_target) / (2.0 * a + b)
631
+ transform = {
632
+ "kind": "continuous_right_skew",
633
+ "tail": "piecewise",
634
+ "low_anchor": low,
635
+ "body_anchor": body_anchor,
636
+ "mid_anchor": mid_anchor,
637
+ "high_anchor": high_anchor,
638
+ "body_target": body_target,
639
+ "mid_target": float(mid_target),
640
+ }
641
+ else:
642
+ # Light tail: finite quadratic from (body_anchor, body_target)
643
+ # to (arr.max, 1) with zero slope at the right edge. The
644
+ # body anchor sits at min(Tukey, arr.max) so bounded
645
+ # distributions don't end up squished. body_target is
646
+ # derived from the geometry to keep slope continuity at the
647
+ # body anchor.
648
+ body_anchor = min(tukey_upper, arr_max)
649
+ body_extent = body_anchor - low
650
+ max_extent = arr_max - low
651
+ body_target = _light_tail_body_target(body_extent, max_extent)
652
+ transform = {
653
+ "kind": "continuous_right_skew",
654
+ "tail": "finite",
655
+ "low_anchor": low,
656
+ "body_anchor": float(body_anchor),
657
+ "high_anchor": float(arr_max),
658
+ "body_target": float(body_target),
659
+ }
660
+ return {
661
+ "transform": transform,
662
+ "impute_value": _compute_impute_value(series),
663
+ "fit_notes": {
664
+ "bowley": float(bowley),
665
+ "tail_asymmetry": float(tail_asymmetry_right),
666
+ },
667
+ }
668
+
669
+ if is_left_skew:
670
+ tukey_lower = float(q1 - 1.5 * iqr)
671
+ p1_lower = float(np.quantile(arr, 1.0 - _COUNT_HIGH_PERCENTILE))
672
+ high = float(arr.max())
673
+ arr_min = float(arr.min())
674
+ if p1_lower < tukey_lower:
675
+ # Heavy tail: 3-segment piecewise (mirror of right). Magnitude
676
+ # runs from high (mag=0) down through body_anchor (Tukey),
677
+ # mid_anchor (p1), to low_anchor (arr.min, mag=1 on the
678
+ # negated axis). See the right-skew block for the
679
+ # mid_target C¹-smoothness derivation.
680
+ body_anchor = float(tukey_lower)
681
+ mid_anchor = float(p1_lower)
682
+ low_anchor = float(arr_min)
683
+ body_target = float(_CONTINUOUS_BODY_TARGET)
684
+ a = body_anchor - mid_anchor
685
+ b = mid_anchor - low_anchor
686
+ if b <= 0.0:
687
+ mid_target = 1.0
688
+ else:
689
+ mid_target = (2.0 * a + b * body_target) / (2.0 * a + b)
690
+ transform = {
691
+ "kind": "continuous_left_skew",
692
+ "tail": "piecewise",
693
+ "low_anchor": low_anchor,
694
+ "body_anchor": body_anchor,
695
+ "mid_anchor": mid_anchor,
696
+ "high_anchor": high,
697
+ "body_target": body_target,
698
+ "mid_target": float(mid_target),
699
+ }
700
+ else:
701
+ # Light tail (mirror of right-skewed light).
702
+ body_anchor = max(tukey_lower, arr_min)
703
+ body_extent = high - body_anchor
704
+ max_extent = high - arr_min
705
+ body_target = _light_tail_body_target(body_extent, max_extent)
706
+ transform = {
707
+ "kind": "continuous_left_skew",
708
+ "tail": "finite",
709
+ "low_anchor": float(arr_min),
710
+ "body_anchor": float(body_anchor),
711
+ "high_anchor": high,
712
+ "body_target": float(body_target),
713
+ }
714
+ return {
715
+ "transform": transform,
716
+ "impute_value": _compute_impute_value(series),
717
+ "fit_notes": {
718
+ "bowley": float(bowley),
719
+ "tail_asymmetry": float(tail_asymmetry_left),
720
+ },
721
+ }
722
+
723
+ # Centered: delegate to _fit_continuous_centered (already in new
724
+ # shape) and just enrich its fit_notes with the bowley reading that
725
+ # routed us here.
726
+ entry = _fit_continuous_centered(series)
727
+ entry["fit_notes"]["bowley"] = float(bowley)
728
+ # Tail-asymmetry is recorded in the direction the bulk-Bowley
729
+ # already points so it's easy to compare against the threshold
730
+ # when wondering "why did this column not go to right/left?".
731
+ if bowley >= 0:
732
+ entry["fit_notes"]["tail_asymmetry"] = float(tail_asymmetry_right)
733
+ else:
734
+ entry["fit_notes"]["tail_asymmetry"] = float(tail_asymmetry_left)
735
+ return entry
736
+
737
+
738
+ def _fit_count(series: pd.Series) -> dict:
739
+ mode = _mode_value(series)
740
+ if mode == 0.0:
741
+ arr = series.dropna().to_numpy(dtype=float)
742
+ q1, q3 = np.quantile(arr, [0.25, 0.75])
743
+ tukey_high = float(q3 + 1.5 * (q3 - q1))
744
+ p99_high = float(np.quantile(arr, _COUNT_HIGH_PERCENTILE))
745
+ high_anchor = max(tukey_high, p99_high)
746
+ if high_anchor <= 0.0:
747
+ high_anchor = float(arr.max())
748
+ if high_anchor <= 0.0:
749
+ return _fit_constant(series)
750
+
751
+ entry = {
752
+ "transform": {"kind": "count_zero_mode", "high_anchor": float(high_anchor)},
753
+ "impute_value": _compute_impute_value(series),
754
+ }
755
+
756
+ # Flag near-degenerate sparse counts where most rows are 0 and
757
+ # the scaler can only produce a handful of distinct values. Goes
758
+ # into fit_notes — advisory only, never read by the transform.
759
+ scaled = np.clip(arr / high_anchor, 0.0, 1.0)
760
+ n_distinct = int(np.unique(scaled).size)
761
+ mode_fraction = float((arr == 0.0).sum()) / float(arr.size)
762
+ if (
763
+ n_distinct <= _DEGENERATE_DISTINCT_MAX
764
+ and mode_fraction >= _DEGENERATE_MODE_FRACTION_MIN
765
+ ):
766
+ entry["fit_notes"] = {
767
+ "degenerate": True,
768
+ "mode_fraction": mode_fraction,
769
+ "n_distinct_output": n_distinct,
770
+ }
771
+ get_logger().warning(
772
+ "Count column '%s' is near-degenerate "
773
+ "(mode_fraction=%.2f, n_distinct_output=%d). "
774
+ "Output collapses to a handful of values — "
775
+ "consider dropping it or revisiting upstream featurization.",
776
+ series.name,
777
+ mode_fraction,
778
+ n_distinct,
779
+ )
780
+ return entry
781
+
782
+ # Count with a non-zero mode. Linear+clip on each side of the mode,
783
+ # with independent upper and lower anchors. No tanh — distinct
784
+ # integer counts in the body and moderate tail keep distinct
785
+ # outputs, which is what users expect from count data.
786
+ #
787
+ # Anchor rule: take the wider of (Tukey whisker, p99) on each side
788
+ # to keep the original outlier-robust property, then cap at the
789
+ # observed data extent so the anchor never reaches past where real
790
+ # data lives — fixes bounded distributions (uniform integers) whose
791
+ # Tukey extends 1.5·IQR past Q3 even though no real data does. The
792
+ # cap doesn't bind for heavy-tail data (where arr.max ≥ Tukey/p99),
793
+ # so outliers past the anchor still clip cleanly.
794
+ arr = series.dropna().to_numpy(dtype=float)
795
+ q1, q3 = np.quantile(arr, [0.25, 0.75])
796
+ iqr = float(q3 - q1)
797
+ tukey_upper = float(q3 + 1.5 * iqr)
798
+ tukey_lower = float(q1 - 1.5 * iqr)
799
+ p99_upper = float(np.quantile(arr, _COUNT_HIGH_PERCENTILE))
800
+ p1_lower = float(np.quantile(arr, 1.0 - _COUNT_HIGH_PERCENTILE))
801
+ arr_max = float(arr.max())
802
+ arr_min = float(arr.min())
803
+ high_anchor = min(max(tukey_upper, p99_upper), arr_max)
804
+ low_anchor = max(min(tukey_lower, p1_lower), arr_min)
805
+ upper_span = max(high_anchor - mode, 0.0)
806
+ lower_span = max(mode - low_anchor, 0.0)
807
+ if upper_span == 0.0 and lower_span == 0.0:
808
+ # Mode is pressed against both edges of the data; nothing to scale.
809
+ return _fit_constant(series)
810
+
811
+ # Soft cap on per-side extent ratio. Distance metrics expect
812
+ # commensurable spread per side; when one side's raw extent is far
813
+ # wider than the other (right- or left-tailed count), without the
814
+ # cap the wider side's distinct discrete values get squashed into
815
+ # the narrow scaled body. Capping the wider extent at
816
+ # ``_COUNT_SHIFTED_EXTENT_RATIO_MAX × the narrower`` sends the
817
+ # excess raw range into the slope-continuous ``tanh`` tail past
818
+ # the anchor, where distinct counts spread distinctly in
819
+ # ``(±0.5, ±1)``. The narrower side is unchanged.
820
+ extent_capped = False
821
+ extent_ratio: Optional[float] = None
822
+ if upper_span > 0.0 and lower_span > 0.0:
823
+ extent_ratio = max(upper_span, lower_span) / min(upper_span, lower_span)
824
+ if extent_ratio > _COUNT_SHIFTED_EXTENT_RATIO_MAX:
825
+ cap = _COUNT_SHIFTED_EXTENT_RATIO_MAX * min(upper_span, lower_span)
826
+ if upper_span > cap:
827
+ upper_span = cap
828
+ high_anchor = mode + upper_span
829
+ if lower_span > cap:
830
+ lower_span = cap
831
+ low_anchor = mode - lower_span
832
+ extent_capped = True
833
+
834
+ entry: dict = {
835
+ "transform": {
836
+ "kind": "count_shifted",
837
+ "center": float(mode),
838
+ "low_anchor": float(low_anchor),
839
+ "high_anchor": float(high_anchor),
840
+ "body_target": float(_CENTERED_BODY_TARGET),
841
+ },
842
+ "impute_value": _compute_impute_value(series),
843
+ }
844
+ if extent_ratio is not None:
845
+ entry["fit_notes"] = {
846
+ "extent_ratio": float(extent_ratio),
847
+ "extent_capped": bool(extent_capped),
848
+ }
849
+ return entry
850
+
851
+
852
+ # ---------------------------------------------------------------------------
853
+ # Apply — per-type transform implementations
854
+ # ---------------------------------------------------------------------------
855
+
856
+
857
+ def _to_output_array(float_arr: np.ndarray, output_dtype: str) -> np.ndarray:
858
+ if output_dtype == "float32":
859
+ return float_arr.astype(np.float32)
860
+ if output_dtype == "int8":
861
+ return _quantize_to_int8(float_arr)
862
+ raise EosframesError(f"Unsupported output_dtype '{output_dtype}'.")
863
+
864
+
865
+ def _quantize_to_int8(float_arr: np.ndarray) -> np.ndarray:
866
+ """Trivially map float ``[-1, 1]`` linearly to int8 ``[-127, 127]``.
867
+
868
+ The mapping is the same for every column: ``int8 = round(x · 127)``,
869
+ clipped to ``[-127, 127]``. Columns whose output region is one-sided
870
+ (binary, right- / left-skewed, count mode 0) naturally inhabit only
871
+ the matching half of the int8 range — a binary 0 stays 0, a binary
872
+ 1 becomes 127. The sentinel ``-128`` is reserved for NaN.
873
+ """
874
+ out = np.full(float_arr.shape, _INT8_NAN_SENTINEL, dtype=np.int8)
875
+ mask = ~np.isnan(float_arr)
876
+ if not mask.any():
877
+ return out
878
+ scaled = np.round(float_arr[mask] * _INT8_MAX_VAL)
879
+ scaled = np.clip(scaled, -_INT8_MAX_VAL, _INT8_MAX_VAL)
880
+ out[mask] = scaled.astype(np.int8)
881
+ return out
882
+
883
+
884
+ def _apply_constant(
885
+ series: pd.Series, _transform: dict, output_dtype: str
886
+ ) -> np.ndarray:
887
+ arr = series.to_numpy(dtype=float)
888
+ out = np.where(np.isnan(arr), np.nan, 0.0)
889
+ return _to_output_array(out, output_dtype)
890
+
891
+
892
+ def _apply_binary(
893
+ series: pd.Series, transform: dict, output_dtype: str
894
+ ) -> np.ndarray:
895
+ low = float(transform["low"])
896
+ high = float(transform["high"])
897
+ arr = series.to_numpy(dtype=float)
898
+ out = np.full(arr.shape, np.nan, dtype=float)
899
+ non_nan = ~np.isnan(arr)
900
+ # Snap every non-NaN value to whichever of {low, high} is closer; ties
901
+ # break to low (a value exactly at the midpoint maps to 0). NaN passes
902
+ # through as NaN (float32) or the int8 sentinel.
903
+ closer_to_high = np.abs(arr - high) < np.abs(arr - low)
904
+ out[non_nan & closer_to_high] = 1.0
905
+ out[non_nan & ~closer_to_high] = 0.0
906
+ return _to_output_array(out, output_dtype)
907
+
908
+
909
+ def _apply_count_zero_mode(
910
+ series: pd.Series, transform: dict, output_dtype: str
911
+ ) -> np.ndarray:
912
+ high_anchor = float(transform["high_anchor"])
913
+ arr = series.to_numpy(dtype=float)
914
+ out = np.where(np.isnan(arr), np.nan, np.clip(arr / high_anchor, 0.0, 1.0))
915
+ return _to_output_array(out, output_dtype)
916
+
917
+
918
+ def _apply_count_shifted(
919
+ series: pd.Series, transform: dict, output_dtype: str
920
+ ) -> np.ndarray:
921
+ """Linear body to ``±body_target`` on each side, ``tanh`` tail toward ``±1``.
922
+
923
+ Mirrors :func:`_apply_continuous_centered`: per-side linear body
924
+ inside ``[mode, anchor]`` (mapping to ``[0, ±body_target]``) plus a
925
+ slope-continuous ``tanh`` asymptote past each anchor. With the
926
+ default ``_CENTERED_BODY_TARGET = 0.5`` the bulk of count data sits
927
+ in ``[-0.5, 0.5]`` and extreme counts stretch asymptotically toward
928
+ ``±1`` rather than piling on a flat clip plateau — distinct
929
+ outliers always get distinct outputs.
930
+ """
931
+ center = float(transform["center"])
932
+ high = float(transform["high_anchor"])
933
+ low = float(transform["low_anchor"])
934
+ body_target = float(transform["body_target"])
935
+ upper_extent = high - center
936
+ lower_extent = center - low
937
+ arr = series.to_numpy(dtype=float)
938
+ nan_mask = np.isnan(arr)
939
+ above = ~nan_mask & (arr >= center)
940
+ below = ~nan_mask & ~above
941
+ out = np.full(arr.shape, np.nan, dtype=float)
942
+
943
+ if above.any() and upper_extent > 0:
944
+ mag = arr[above] - center
945
+ in_body = mag <= upper_extent
946
+ y_up = np.empty_like(mag)
947
+ y_up[in_body] = mag[in_body] / upper_extent * body_target
948
+ if (~in_body).any():
949
+ y_up[~in_body] = _tanh_tail(mag[~in_body], upper_extent, body_target)
950
+ out[above] = y_up
951
+ elif above.any():
952
+ out[above] = 0.0
953
+
954
+ if below.any() and lower_extent > 0:
955
+ mag = center - arr[below]
956
+ in_body = mag <= lower_extent
957
+ y_lo = np.empty_like(mag)
958
+ y_lo[in_body] = mag[in_body] / lower_extent * body_target
959
+ if (~in_body).any():
960
+ y_lo[~in_body] = _tanh_tail(mag[~in_body], lower_extent, body_target)
961
+ out[below] = -y_lo
962
+ elif below.any():
963
+ out[below] = 0.0
964
+ return _to_output_array(out, output_dtype)
965
+
966
+
967
+ def _apply_continuous_right(
968
+ series: pd.Series, transform: dict, output_dtype: str
969
+ ) -> np.ndarray:
970
+ """Linear body + tail dispatch (3-segment piecewise or finite quadratic).
971
+
972
+ Heavy tail (``tail == "piecewise"``): 3-segment design. Linear
973
+ body to ``(body_anchor, body_target)``, then a slower linear
974
+ "middle" segment to ``(mid_anchor, mid_target)``, then a
975
+ finite quadratic to ``(high_anchor, 1)`` with C¹ slope at the
976
+ middle→tail join. Outliers spread visibly across
977
+ ``[body_target, 1]`` instead of saturating asymptotically.
978
+
979
+ Light tail (``tail == "finite"``): quadratic from
980
+ ``(body_anchor, body_target)`` to ``(high_anchor, 1)`` with
981
+ slope-continuous join at the body and zero slope at
982
+ ``high_anchor``. ``body_target`` is derived from the geometry so
983
+ the quadratic reaches exactly 1 at ``high_anchor``; values past
984
+ ``high_anchor`` clip to 1 (only happens for OOD inputs).
985
+ """
986
+ low = float(transform["low_anchor"])
987
+ body = float(transform["body_anchor"])
988
+ body_target = float(transform["body_target"])
989
+ tail_kind = transform.get("tail", "piecewise")
990
+ body_span = body - low
991
+ arr = series.to_numpy(dtype=float)
992
+ nan = np.isnan(arr)
993
+ out = np.full(arr.shape, np.nan, dtype=float)
994
+
995
+ body_mask = ~nan & (arr <= body)
996
+ tail_mask = ~nan & (arr > body)
997
+
998
+ if body_span > 0:
999
+ out[body_mask] = np.clip(
1000
+ (arr[body_mask] - low) / body_span * body_target,
1001
+ 0.0,
1002
+ body_target,
1003
+ )
1004
+ else:
1005
+ out[body_mask] = 0.0
1006
+
1007
+ if tail_mask.any():
1008
+ mag = arr[tail_mask] - low # distance from low_anchor, ≥ body_span
1009
+ if tail_kind == "piecewise":
1010
+ mid = float(transform["mid_anchor"])
1011
+ high = float(transform["high_anchor"])
1012
+ mid_target = float(transform["mid_target"])
1013
+ mid_end = mid - low
1014
+ max_end = high - low
1015
+ seg2 = mag <= mid_end
1016
+ seg3 = (mag > mid_end) & (mag <= max_end)
1017
+ over = mag > max_end
1018
+ y = np.empty_like(mag)
1019
+ mid_span = mid_end - body_span
1020
+ if mid_span > 0 and seg2.any():
1021
+ y[seg2] = body_target + (mag[seg2] - body_span) / mid_span * (
1022
+ mid_target - body_target
1023
+ )
1024
+ else:
1025
+ y[seg2] = body_target
1026
+ if seg3.any():
1027
+ y[seg3] = _quadratic_tail(mag[seg3], mid_end, mid_target, max_end)
1028
+ y[over] = 1.0
1029
+ out[tail_mask] = y
1030
+ else: # "finite"
1031
+ high = float(transform["high_anchor"])
1032
+ max_extent = high - low
1033
+ out[tail_mask] = _quadratic_tail(mag, body_span, body_target, max_extent)
1034
+
1035
+ # Cubic Hermite blend across the body→middle junction (piecewise
1036
+ # tail only). Smooths the steep-to-gentle slope change so the
1037
+ # scaled histogram doesn't show a sharp "double decay" hump at
1038
+ # body_target. The cubic interpolates between (body-δ, body slope)
1039
+ # and (body+δ, mid slope) with prescribed endpoint values; for the
1040
+ # slope ratios in heavy-tail right skew the cubic is monotone
1041
+ # (Fritsch-Carlson condition holds for `body_target ≤ 1` cases).
1042
+ if tail_kind == "piecewise" and body_span > 0:
1043
+ mid = float(transform["mid_anchor"])
1044
+ high = float(transform["high_anchor"])
1045
+ mid_target = float(transform["mid_target"])
1046
+ mid_end = mid - low
1047
+ mid_span = mid_end - body_span
1048
+ if mid_span > 0:
1049
+ delta = _PIECEWISE_BLEND_FRACTION * min(body_span, mid_span)
1050
+ blend_lo = low + body_span - delta
1051
+ blend_hi = low + body_span + delta
1052
+ blend_mask = ~nan & (arr > blend_lo) & (arr < blend_hi)
1053
+ if blend_mask.any():
1054
+ mag_b = arr[blend_mask] - low
1055
+ # Endpoint values & slopes for the cubic.
1056
+ p0 = (body_span - delta) / body_span * body_target
1057
+ p1 = body_target + delta / mid_span * (mid_target - body_target)
1058
+ m_body = body_target / body_span
1059
+ m_mid = (mid_target - body_target) / mid_span
1060
+ dx = 2.0 * delta
1061
+ t = (mag_b - (body_span - delta)) / dx
1062
+ t2 = t * t
1063
+ t3 = t2 * t
1064
+ h00 = 2.0 * t3 - 3.0 * t2 + 1.0
1065
+ h10 = t3 - 2.0 * t2 + t
1066
+ h01 = -2.0 * t3 + 3.0 * t2
1067
+ h11 = t3 - t2
1068
+ out[blend_mask] = (
1069
+ h00 * p0 + h10 * dx * m_body + h01 * p1 + h11 * dx * m_mid
1070
+ )
1071
+
1072
+ return _to_output_array(out, output_dtype)
1073
+
1074
+
1075
+ def _apply_continuous_left(
1076
+ series: pd.Series, transform: dict, output_dtype: str
1077
+ ) -> np.ndarray:
1078
+ """Mirror of right-skew. Output region is ``[-1, 0]``."""
1079
+ high = float(transform["high_anchor"])
1080
+ body = float(transform["body_anchor"])
1081
+ body_target = float(transform["body_target"])
1082
+ tail_kind = transform.get("tail", "piecewise")
1083
+ body_span = high - body
1084
+ arr = series.to_numpy(dtype=float)
1085
+ nan = np.isnan(arr)
1086
+ out = np.full(arr.shape, np.nan, dtype=float)
1087
+
1088
+ body_mask = ~nan & (arr >= body)
1089
+ tail_mask = ~nan & (arr < body)
1090
+
1091
+ if body_span > 0:
1092
+ out[body_mask] = np.clip(
1093
+ -body_target + (arr[body_mask] - body) / body_span * body_target,
1094
+ -body_target,
1095
+ 0.0,
1096
+ )
1097
+ else:
1098
+ out[body_mask] = 0.0
1099
+
1100
+ if tail_mask.any():
1101
+ mag = high - arr[tail_mask] # distance toward the left tail, ≥ body_span
1102
+ if tail_kind == "piecewise":
1103
+ mid = float(transform["mid_anchor"])
1104
+ low = float(transform["low_anchor"])
1105
+ mid_target = float(transform["mid_target"])
1106
+ mid_end = high - mid
1107
+ max_end = high - low
1108
+ seg2 = mag <= mid_end
1109
+ seg3 = (mag > mid_end) & (mag <= max_end)
1110
+ over = mag > max_end
1111
+ y = np.empty_like(mag)
1112
+ mid_span = mid_end - body_span
1113
+ if mid_span > 0 and seg2.any():
1114
+ y[seg2] = body_target + (mag[seg2] - body_span) / mid_span * (
1115
+ mid_target - body_target
1116
+ )
1117
+ else:
1118
+ y[seg2] = body_target
1119
+ if seg3.any():
1120
+ y[seg3] = _quadratic_tail(mag[seg3], mid_end, mid_target, max_end)
1121
+ y[over] = 1.0
1122
+ out[tail_mask] = -y
1123
+ else: # "finite"
1124
+ low = float(transform["low_anchor"])
1125
+ max_extent = high - low
1126
+ out[tail_mask] = -_quadratic_tail(
1127
+ mag, body_span, body_target, max_extent
1128
+ )
1129
+
1130
+ # Cubic Hermite blend across the body→middle junction (mirror of
1131
+ # the right-skew blend).
1132
+ if tail_kind == "piecewise" and body_span > 0:
1133
+ mid = float(transform["mid_anchor"])
1134
+ low = float(transform["low_anchor"])
1135
+ mid_target = float(transform["mid_target"])
1136
+ mid_end = high - mid
1137
+ mid_span = mid_end - body_span
1138
+ if mid_span > 0:
1139
+ delta = _PIECEWISE_BLEND_FRACTION * min(body_span, mid_span)
1140
+ blend_lo = body - delta
1141
+ blend_hi = body + delta
1142
+ blend_mask = ~nan & (arr > blend_lo) & (arr < blend_hi)
1143
+ if blend_mask.any():
1144
+ mag_b = high - arr[blend_mask]
1145
+ p0 = (body_span - delta) / body_span * body_target
1146
+ p1 = body_target + delta / mid_span * (mid_target - body_target)
1147
+ m_body = body_target / body_span
1148
+ m_mid = (mid_target - body_target) / mid_span
1149
+ dx = 2.0 * delta
1150
+ t = (mag_b - (body_span - delta)) / dx
1151
+ t2 = t * t
1152
+ t3 = t2 * t
1153
+ h00 = 2.0 * t3 - 3.0 * t2 + 1.0
1154
+ h10 = t3 - 2.0 * t2 + t
1155
+ h01 = -2.0 * t3 + 3.0 * t2
1156
+ h11 = t3 - t2
1157
+ out[blend_mask] = -(
1158
+ h00 * p0 + h10 * dx * m_body + h01 * p1 + h11 * dx * m_mid
1159
+ )
1160
+
1161
+ return _to_output_array(out, output_dtype)
1162
+
1163
+
1164
+ def _apply_continuous_centered(
1165
+ series: pd.Series, transform: dict, output_dtype: str
1166
+ ) -> np.ndarray:
1167
+ """Per-side linear body + finite-reach quadratic tail.
1168
+
1169
+ For each side independently: linear from 0 at the centre to
1170
+ ``±body_target`` at ``center ± *_body_extent``, then
1171
+ :func:`_quadratic_tail` to ``±1`` at ``center ± *_effective_max``,
1172
+ where ``effective_max = max(*_max_extent,
1173
+ _CENTERED_MAX_EXTENT_RATIO · body_extent)``. ``body_target`` is
1174
+ derived per side from ``(body_extent, effective_max)`` and
1175
+ **capped at** :data:`_CENTERED_BODY_TARGET` (``0.5``) so the bulk
1176
+ of the density visually lives in ``[-0.5, 0.5]``, matching
1177
+ ``count_shifted``. The dual cap (``0.5``) and floor (``3·b``) are
1178
+ mathematically locked: the body→tail join is C¹-smooth precisely
1179
+ at ``max_extent = 3·body_extent`` for ``body_target = 0.5``, so
1180
+ flooring the effective max at ``3·b`` whenever the data doesn't
1181
+ naturally provide it both eliminates the slope kink and prevents
1182
+ bounded distributions (U-shape, narrow bimodal, truncated normal)
1183
+ from spraying their few near-edge points across the entire tail
1184
+ region. Inputs past ``effective_max`` clip to exactly ``±1``.
1185
+ """
1186
+ center = float(transform["center"])
1187
+ upper_body = float(transform["upper_body_extent"])
1188
+ lower_body = float(transform["lower_body_extent"])
1189
+ upper_max = float(transform["upper_max_extent"])
1190
+ lower_max = float(transform["lower_max_extent"])
1191
+
1192
+ arr = series.to_numpy(dtype=float)
1193
+ nan = np.isnan(arr)
1194
+ out = np.full(arr.shape, np.nan, dtype=float)
1195
+
1196
+ def _side(mask: np.ndarray, body_extent: float, max_extent: float, sign: float) -> None:
1197
+ if not mask.any():
1198
+ return
1199
+ if body_extent <= 0:
1200
+ out[mask] = 0.0
1201
+ return
1202
+ effective_max = max(max_extent, _CENTERED_MAX_EXTENT_RATIO * body_extent)
1203
+ body_target = min(
1204
+ _light_tail_body_target(body_extent, effective_max),
1205
+ _CENTERED_BODY_TARGET,
1206
+ )
1207
+ mag = sign * (arr[mask] - center)
1208
+ body_part = mag <= body_extent
1209
+ tail_part = ~body_part
1210
+ y = np.empty_like(mag)
1211
+ y[body_part] = mag[body_part] / body_extent * body_target
1212
+ if tail_part.any():
1213
+ y[tail_part] = _quadratic_tail(
1214
+ mag[tail_part], body_extent, body_target, effective_max
1215
+ )
1216
+ out[mask] = sign * y
1217
+
1218
+ _side(~nan & (arr >= center), upper_body, upper_max, +1.0)
1219
+ _side(~nan & (arr < center), lower_body, lower_max, -1.0)
1220
+ return _to_output_array(out, output_dtype)
1221
+
1222
+
1223
+ # Single-string dispatch: each ``kind`` maps to the apply function that
1224
+ # reads its specific ``transform`` payload. Keep in sync with
1225
+ # ``_OUTPUT_REGIONS`` — the two tables share the same key set.
1226
+ _APPLY_DISPATCH: Dict[str, Callable[[pd.Series, dict, str], np.ndarray]] = {
1227
+ "constant": _apply_constant,
1228
+ "binary": _apply_binary,
1229
+ "count_zero_mode": _apply_count_zero_mode,
1230
+ "count_shifted": _apply_count_shifted,
1231
+ "continuous_right_skew": _apply_continuous_right,
1232
+ "continuous_left_skew": _apply_continuous_left,
1233
+ "continuous_centered": _apply_continuous_centered,
1234
+ }
1235
+
1236
+
1237
+ def _dispatch_apply(
1238
+ series: pd.Series, transform: dict, output_dtype: str
1239
+ ) -> np.ndarray:
1240
+ kind = transform["kind"]
1241
+ apply_fn = _APPLY_DISPATCH.get(kind)
1242
+ if apply_fn is None:
1243
+ raise EosframesError(f"Unknown transform kind '{kind}'.")
1244
+ return apply_fn(series, transform, output_dtype)
1245
+
1246
+
1247
+ # ---------------------------------------------------------------------------
1248
+ # Low-level DataFrame API
1249
+ # ---------------------------------------------------------------------------
1250
+
1251
+
1252
+ def fit(df: pd.DataFrame) -> dict:
1253
+ """Fit a type-aware robust scaler on the numeric feature columns.
1254
+
1255
+ Each numeric feature column is classified into one of the seven
1256
+ transform kinds (``constant`` / ``binary`` / ``count_zero_mode`` /
1257
+ ``count_shifted`` / ``continuous_right_skew`` /
1258
+ ``continuous_left_skew`` / ``continuous_centered``) and a per-column
1259
+ entry is recorded. Every numeric column is fitted; all-NaN columns
1260
+ fall back to ``kind: "constant"``. The ``key`` and ``input``
1261
+ columns and any non-numeric feature columns are ignored entirely.
1262
+
1263
+ Output dtype is a **transform-time** choice, not a fit-time one —
1264
+ pass ``output_dtype`` to :func:`transform` (or ``--quantize`` to
1265
+ the CLI).
1266
+
1267
+ Parameters
1268
+ ----------
1269
+ df : pandas.DataFrame
1270
+ Input frame. ``key`` and ``input`` columns are ignored.
1271
+
1272
+ Returns
1273
+ -------
1274
+ dict
1275
+ Dtype-agnostic parameters with keys:
1276
+
1277
+ * ``method`` — always ``"robust_typed"``.
1278
+ * ``columns`` — ``{column_name: entry}`` for every fitted
1279
+ column, in fit-time order. Each ``entry`` has ``transform``
1280
+ (the kind + transform-time params), ``impute_value`` (median
1281
+ fill for ``impute=True``), and an optional ``fit_notes``
1282
+ (provenance diagnostics, never read at transform time).
1283
+
1284
+ Raises
1285
+ ------
1286
+ EosframesError
1287
+ If no numeric columns exist.
1288
+ """
1289
+ logger = get_logger()
1290
+ feature_cols = [c for c in df.columns if c not in _META_COLS]
1291
+ numeric_cols = [c for c in feature_cols if pd.api.types.is_numeric_dtype(df[c])]
1292
+
1293
+ if not numeric_cols:
1294
+ raise EosframesError("No numeric feature columns found to fit the scaler.")
1295
+
1296
+ columns: dict = {}
1297
+ kind_counts: dict = {}
1298
+
1299
+ for col in numeric_cols:
1300
+ series = df[col]
1301
+
1302
+ if series.dropna().empty:
1303
+ # All-NaN column: fit as a constant. The dispatch at
1304
+ # transform time maps non-NaN inputs to 0 and propagates NaN.
1305
+ entry = _fit_constant(series)
1306
+ else:
1307
+ type_ = _classify_type(series)
1308
+ if type_ == "constant":
1309
+ entry = _fit_constant(series)
1310
+ elif type_ == "binary":
1311
+ entry = _fit_binary(series)
1312
+ elif type_ == "count":
1313
+ entry = _fit_count(series)
1314
+ else:
1315
+ try:
1316
+ entry = _fit_continuous(series)
1317
+ except EosframesError:
1318
+ # Column slipped past _classify_type but has no usable scale.
1319
+ entry = _fit_constant(series)
1320
+
1321
+ columns[col] = entry
1322
+ kind = entry["transform"]["kind"]
1323
+ kind_counts[kind] = kind_counts.get(kind, 0) + 1
1324
+
1325
+ kind_breakdown = ", ".join(
1326
+ f"{kind}={count}" for kind, count in sorted(kind_counts.items())
1327
+ )
1328
+ logger.info(
1329
+ "Fitted %d / %d numeric columns (%s)",
1330
+ len(numeric_cols),
1331
+ len(numeric_cols),
1332
+ kind_breakdown,
1333
+ )
1334
+
1335
+ return {"method": _METHOD_NAME, "columns": columns}
1336
+
1337
+
1338
+ def transform(
1339
+ df: pd.DataFrame,
1340
+ params: dict,
1341
+ output_dtype: str = _DEFAULT_OUTPUT_DTYPE,
1342
+ impute: bool = False,
1343
+ ) -> pd.DataFrame:
1344
+ """Apply a fitted scaler to a DataFrame.
1345
+
1346
+ The ``key`` / ``input`` columns pass through unchanged. The fitted
1347
+ feature columns in *df* must match exactly — same set and same
1348
+ order — the ``feature_columns`` recorded in *params*.
1349
+
1350
+ Parameters
1351
+ ----------
1352
+ df : pandas.DataFrame
1353
+ Frame to transform. Must have the same numeric feature columns
1354
+ as the frame the scaler was fitted on.
1355
+ params : dict
1356
+ Dtype-agnostic parameters as returned by :func:`fit` or loaded
1357
+ from a scaler JSON.
1358
+ output_dtype : {"float32", "int8"}, default ``"float32"``
1359
+ ``"float32"`` preserves NaN for missing values. ``"int8"``
1360
+ quantizes scaled values to ``[-127, 127]`` with sentinel
1361
+ ``-128`` for missing — useful for compact storage of fingerprint-
1362
+ like outputs.
1363
+ impute : bool, default ``False``
1364
+ When ``True``, replace every input NaN with the column's
1365
+ recorded ``impute_value`` *before* dispatch. The output column
1366
+ will have no NaN entries (and, under ``output_dtype="int8"``,
1367
+ no ``-128`` sentinels). The substituted value is the median of
1368
+ the column's training data, rounded to ``int`` if the column
1369
+ was integer-valued at fit time.
1370
+
1371
+ Returns
1372
+ -------
1373
+ pandas.DataFrame
1374
+ Copy of *df* with scaled feature columns in the requested
1375
+ dtype. The returned frame does **not** carry ``model_id`` /
1376
+ ``version`` attributes — re-attach them before writing.
1377
+
1378
+ Raises
1379
+ ------
1380
+ EosframesError
1381
+ On column mismatch, invalid ``output_dtype``, or unknown
1382
+ column type in *params*.
1383
+ """
1384
+ if output_dtype not in _VALID_OUTPUT_DTYPES:
1385
+ raise EosframesError(
1386
+ f"Unknown output_dtype '{output_dtype}'. "
1387
+ f"Supported: {_VALID_OUTPUT_DTYPES}"
1388
+ )
1389
+
1390
+ expected_feature_cols = list(params["columns"].keys())
1391
+ df_feature_cols = [c for c in df.columns if c not in _META_COLS]
1392
+
1393
+ if df_feature_cols != expected_feature_cols:
1394
+ raise EosframesError(
1395
+ f"Column mismatch: input has feature columns {df_feature_cols} "
1396
+ f"but transformer was fitted on {expected_feature_cols}."
1397
+ )
1398
+
1399
+ columns = params["columns"]
1400
+ result = df.copy()
1401
+ for col in expected_feature_cols:
1402
+ entry = columns[col]
1403
+ series = df[col]
1404
+ if impute:
1405
+ series = series.fillna(entry.get("impute_value", 0.0))
1406
+ result[col] = _dispatch_apply(series, entry["transform"], output_dtype)
1407
+ return result
1408
+
1409
+
1410
+ # ---------------------------------------------------------------------------
1411
+ # File-level API
1412
+ # ---------------------------------------------------------------------------
1413
+
1414
+
1415
+ def _values_dtype_for(output_dtype: str) -> np.dtype:
1416
+ if output_dtype == "int8":
1417
+ return np.int8
1418
+ return np.float32
1419
+
1420
+
1421
+ def _write_df(
1422
+ df: pd.DataFrame, output_path: str, values_dtype: np.dtype = np.float32
1423
+ ) -> None:
1424
+ """Write a DataFrame to CSV or H5, bypassing the naming convention check.
1425
+
1426
+ *values_dtype* controls the H5 ``values`` dataset dtype; CSV ignores it.
1427
+ """
1428
+ ext = os.path.splitext(output_path)[1].lower()
1429
+ if ext == ".csv":
1430
+ df.to_csv(output_path, index=False)
1431
+ elif ext == ".h5":
1432
+ feat_cols = [c for c in df.columns if c not in _META_COLS]
1433
+ with h5py.File(output_path, "w") as f:
1434
+ dt = h5py.string_dtype(encoding="utf-8")
1435
+ if "key" in df.columns:
1436
+ f.create_dataset("key", data=df["key"].astype(str).tolist(), dtype=dt)
1437
+ if "input" in df.columns:
1438
+ f.create_dataset(
1439
+ "input", data=df["input"].astype(str).tolist(), dtype=dt
1440
+ )
1441
+ f.create_dataset("features", data=feat_cols, dtype=dt)
1442
+ f.create_dataset(
1443
+ "values", data=df[feat_cols].values, dtype=values_dtype
1444
+ )
1445
+ else:
1446
+ raise EosframesError(f"Unsupported output format '{ext}'. Expected .csv or .h5")
1447
+
1448
+
1449
+ def fit_file(
1450
+ input_path: str,
1451
+ scaler_path: str,
1452
+ output_path: Optional[str] = None,
1453
+ output_dtype: str = _DEFAULT_OUTPUT_DTYPE,
1454
+ impute: bool = False,
1455
+ ) -> str:
1456
+ """Fit a scaler on an Ersilia output file and save the parameters.
1457
+
1458
+ The scaler JSON written to *scaler_path* is dtype-agnostic — see
1459
+ :func:`fit` for the parameter set. When *output_path* is provided
1460
+ the scaled data is also written immediately (fit-then-transform in
1461
+ one call), and *output_dtype* selects the dtype of that inline
1462
+ output. The dtype is **not** recorded in the scaler JSON; later
1463
+ calls to :func:`transform_file` choose the dtype independently.
1464
+
1465
+ The scaler filename's encoded model ID and version must match the
1466
+ input file's. The transformer JSON records ``eosframes_version``
1467
+ (the running package version, via ``importlib.metadata``) and
1468
+ ``method`` (``"robust_typed"``); :func:`transform_file` rejects on
1469
+ any mismatch with the current ``eosframes.__version__`` so a
1470
+ package release automatically forces a re-fit.
1471
+
1472
+ Parameters
1473
+ ----------
1474
+ input_path : str
1475
+ Input CSV or H5 file. Must follow the Ersilia naming
1476
+ convention.
1477
+ scaler_path : str
1478
+ Path where the JSON parameter file will be written. Must not
1479
+ already exist. Must follow
1480
+ ``[prefix_]<model_id>_<version>_transformer.json`` with model
1481
+ ID and version matching *input_path*.
1482
+ output_path : str, optional
1483
+ If provided, also write the scaled data here (fit-transform).
1484
+ The file must not already exist; its extension determines
1485
+ whether CSV or H5 is written.
1486
+ output_dtype : {"float32", "int8"}, default ``"float32"``
1487
+ Only used for the inline transform when *output_path* is
1488
+ given. ``"int8"`` quantizes the scaled values into
1489
+ ``[-127, 127]`` with sentinel ``-128`` for missing.
1490
+
1491
+ Returns
1492
+ -------
1493
+ str
1494
+ Absolute path of the saved scaler JSON file.
1495
+
1496
+ Raises
1497
+ ------
1498
+ EosframesError
1499
+ On naming convention violations, pre-existing files,
1500
+ model-ID / version mismatch between scaler and input, no
1501
+ numeric columns to fit, or invalid ``output_dtype``.
1502
+ """
1503
+ logger = get_logger()
1504
+
1505
+ if not is_valid_name(input_path):
1506
+ raise EosframesError(
1507
+ f"'{input_path}' does not follow the naming convention. "
1508
+ "Expected: <model_id>_<version>.<ext>"
1509
+ )
1510
+
1511
+ if not is_valid_transformer_name(scaler_path):
1512
+ raise EosframesError(
1513
+ f"'{scaler_path}' does not follow the scaler naming convention. "
1514
+ "Expected: [prefix_]<model_id>_<version>_transformer.json"
1515
+ )
1516
+
1517
+ parsed = parse_name(input_path)
1518
+ scaler_parsed = parse_transformer_name(scaler_path)
1519
+
1520
+ if scaler_parsed["model_id"] != parsed["model_id"]:
1521
+ raise EosframesError(
1522
+ f"Scaler model ID '{scaler_parsed['model_id']}' does not match "
1523
+ f"input model ID '{parsed['model_id']}'. "
1524
+ "Scaler filename must encode the same model ID as the input file."
1525
+ )
1526
+ if scaler_parsed["version"] != parsed["version"]:
1527
+ raise EosframesError(
1528
+ f"Scaler version '{scaler_parsed['version']}' does not match "
1529
+ f"input version '{parsed['version']}'. "
1530
+ "Scaler filename must encode the same version as the input file."
1531
+ )
1532
+
1533
+ if os.path.exists(scaler_path):
1534
+ raise EosframesError(
1535
+ f"Scaler file '{scaler_path}' already exists. Remove it first."
1536
+ )
1537
+
1538
+ if output_path is not None and os.path.exists(output_path):
1539
+ raise EosframesError(
1540
+ f"Output file '{output_path}' already exists. Remove it first."
1541
+ )
1542
+
1543
+ from .ops import _read_file
1544
+
1545
+ df = _read_file(input_path)
1546
+
1547
+ if output_dtype not in _VALID_OUTPUT_DTYPES:
1548
+ raise EosframesError(
1549
+ f"Unknown output_dtype '{output_dtype}'. "
1550
+ f"Supported: {_VALID_OUTPUT_DTYPES}"
1551
+ )
1552
+
1553
+ mode_label = "fit + transform" if output_path is not None else "fit only"
1554
+ logger.info(
1555
+ "%s — reading %d rows from %s", mode_label, len(df), input_path
1556
+ )
1557
+ fitted = fit(df)
1558
+
1559
+ transformer = {
1560
+ "eosframes_version": _PACKAGE_VERSION,
1561
+ "method": fitted["method"],
1562
+ "model_id": parsed["model_id"],
1563
+ "model_version": parsed["version"],
1564
+ "fitted_at": datetime.now().isoformat(timespec="seconds"),
1565
+ "n_rows": len(df),
1566
+ "columns": fitted["columns"],
1567
+ }
1568
+ with open(scaler_path, "w") as fh:
1569
+ json.dump(transformer, fh, indent=2)
1570
+ logger.info("Scaler saved to %s", scaler_path)
1571
+
1572
+ if output_path is not None:
1573
+ logger.info(
1574
+ "Transforming inline (output_dtype=%s, impute=%s)",
1575
+ output_dtype,
1576
+ impute,
1577
+ )
1578
+ scaled_df = transform(
1579
+ df, transformer, output_dtype=output_dtype, impute=impute
1580
+ )
1581
+ _write_df(
1582
+ scaled_df,
1583
+ output_path,
1584
+ values_dtype=_values_dtype_for(output_dtype),
1585
+ )
1586
+ logger.info("Scaled output written to %s", output_path)
1587
+
1588
+ return scaler_path
1589
+
1590
+
1591
+ def transform_file(
1592
+ input_path: str,
1593
+ scaler_path: str,
1594
+ output_path: str,
1595
+ output_dtype: str = _DEFAULT_OUTPUT_DTYPE,
1596
+ impute: bool = False,
1597
+ ) -> str:
1598
+ """Apply a saved scaler to an Ersilia output file.
1599
+
1600
+ The scaler's recorded ``eosframes_version`` must exactly match the
1601
+ running ``eosframes.__version__``, and ``method`` must still be
1602
+ ``"robust_typed"``; any mismatch raises with a clear "re-fit"
1603
+ error. The scaler's recorded ``model_id`` and ``version`` must
1604
+ match the model ID / version encoded in *input_path* — running a
1605
+ scaler against a different model's outputs is never silently
1606
+ allowed.
1607
+
1608
+ Parameters
1609
+ ----------
1610
+ input_path : str
1611
+ Input CSV or H5 file. Must follow the Ersilia naming
1612
+ convention.
1613
+ scaler_path : str
1614
+ Path to a JSON scaler file produced by :func:`fit_file`.
1615
+ output_path : str
1616
+ Where to write the scaled data. Must not already exist; its
1617
+ extension determines whether CSV or H5 is written.
1618
+ output_dtype : {"float32", "int8"}, default ``"float32"``
1619
+ ``"int8"`` quantizes scaled values to ``[-127, 127]`` with
1620
+ sentinel ``-128`` for missing.
1621
+
1622
+ Returns
1623
+ -------
1624
+ str
1625
+ Absolute path of the scaled output file.
1626
+
1627
+ Raises
1628
+ ------
1629
+ EosframesError
1630
+ On naming-convention violations, missing scaler file,
1631
+ pre-existing output path, ``eosframes_version`` or ``method``
1632
+ mismatch on the scaler JSON, model-ID or version mismatch
1633
+ between scaler and input, column mismatch between scaler and
1634
+ input, or invalid ``output_dtype``.
1635
+ """
1636
+ logger = get_logger()
1637
+
1638
+ if not is_valid_name(input_path):
1639
+ raise EosframesError(
1640
+ f"'{input_path}' does not follow the naming convention. "
1641
+ "Expected: <model_id>_<version>.<ext>"
1642
+ )
1643
+
1644
+ if not is_valid_transformer_name(scaler_path):
1645
+ raise EosframesError(
1646
+ f"'{scaler_path}' does not follow the scaler naming convention. "
1647
+ "Expected: [prefix_]<model_id>_<version>_transformer.json"
1648
+ )
1649
+
1650
+ parsed = parse_name(input_path)
1651
+
1652
+ if os.path.exists(output_path):
1653
+ raise EosframesError(
1654
+ f"Output file '{output_path}' already exists. Remove it first."
1655
+ )
1656
+
1657
+ if not os.path.exists(scaler_path):
1658
+ raise EosframesError(f"Scaler file '{scaler_path}' not found.")
1659
+
1660
+ with open(scaler_path) as fh:
1661
+ transformer = json.load(fh)
1662
+
1663
+ scaler_version = transformer.get("eosframes_version") or ""
1664
+ method = transformer.get("method")
1665
+ if _major(scaler_version) != _major(_PACKAGE_VERSION) or method != _METHOD_NAME:
1666
+ raise EosframesError(
1667
+ f"Scaler was fitted with eosframes "
1668
+ f"{scaler_version!r} (method={method!r}) but this is eosframes "
1669
+ f"{_PACKAGE_VERSION!r} (method '{_METHOD_NAME}'). Re-fit the "
1670
+ "scaler — the schema only carries across the same eosframes "
1671
+ "major version."
1672
+ )
1673
+
1674
+ if output_dtype not in _VALID_OUTPUT_DTYPES:
1675
+ raise EosframesError(
1676
+ f"Unknown output_dtype '{output_dtype}'. "
1677
+ f"Supported: {_VALID_OUTPUT_DTYPES}"
1678
+ )
1679
+
1680
+ t_model_id = transformer.get("model_id")
1681
+ t_model_version = transformer.get("model_version")
1682
+ f_model_id = parsed["model_id"]
1683
+ f_model_version = parsed["version"]
1684
+
1685
+ if f_model_id != t_model_id:
1686
+ raise EosframesError(
1687
+ f"Model ID mismatch: file has '{f_model_id}' but scaler "
1688
+ f"was fitted on '{t_model_id}'."
1689
+ )
1690
+ if f_model_version != t_model_version:
1691
+ raise EosframesError(
1692
+ f"Model version mismatch: file has '{f_model_version}' but "
1693
+ f"scaler was fitted on '{t_model_version}'."
1694
+ )
1695
+
1696
+ from .ops import _read_file
1697
+
1698
+ df = _read_file(input_path)
1699
+
1700
+ logger.info(
1701
+ "Applying scaler to %d rows from %s (output_dtype=%s, impute=%s)",
1702
+ len(df),
1703
+ input_path,
1704
+ output_dtype,
1705
+ impute,
1706
+ )
1707
+ scaled_df = transform(
1708
+ df, transformer, output_dtype=output_dtype, impute=impute
1709
+ )
1710
+
1711
+ _write_df(scaled_df, output_path, values_dtype=_values_dtype_for(output_dtype))
1712
+ logger.info("Scaled output written to %s", output_path)
1713
+ return output_path