MaldiDeepKit 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.
@@ -0,0 +1,492 @@
1
+ """1-D Vision Transformer for binned MALDI-TOF spectra.
2
+
3
+ A plain ViT backbone adapted to 1-D spectra:
4
+ non-overlapping patch embedding, learned positional embedding,
5
+ pre-LayerNorm residual blocks with LayerScale and stochastic depth,
6
+ global self-attention in every block, and mean-pool aggregation by
7
+ default (CLS token optional).
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from pathlib import Path
13
+ from typing import Any
14
+
15
+ import numpy as np
16
+ import torch
17
+ from torch import nn
18
+
19
+ from .._blocks import DropPath, PatchEmbed1D
20
+ from ..base.classifier import BaseSpectralClassifier
21
+
22
+
23
+ class MultiHeadSelfAttention(nn.Module):
24
+ """Multi-head self-attention with QK-norm + memory-efficient SDPA.
25
+
26
+ QK-normalization applies a per-head :class:`~torch.nn.LayerNorm`
27
+ to query and key tensors before the scaled-dot-product, bounding
28
+ the softmax denominator regardless of input scale. Always on
29
+ (universal stability improvement with negligible compute overhead).
30
+
31
+ Parameters
32
+ ----------
33
+ dim : int
34
+ Token embedding dimension. Must be divisible by ``num_heads``.
35
+ num_heads : int
36
+ Number of attention heads.
37
+ attention_dropout : float, default=0.0
38
+ Dropout applied inside the attention kernel during training.
39
+ proj_dropout : float, default=0.0
40
+ Dropout applied to the final projection.
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ dim: int,
46
+ num_heads: int,
47
+ attention_dropout: float = 0.0,
48
+ proj_dropout: float = 0.0,
49
+ ) -> None:
50
+ super().__init__()
51
+ if dim % num_heads != 0:
52
+ raise ValueError(f"dim={dim} must be divisible by num_heads={num_heads}.")
53
+ self.num_heads = num_heads
54
+ self.head_dim = dim // num_heads
55
+ self.qkv = nn.Linear(dim, 3 * dim)
56
+ self.q_norm = nn.LayerNorm(self.head_dim)
57
+ self.k_norm = nn.LayerNorm(self.head_dim)
58
+ self.attention_dropout = float(attention_dropout)
59
+ self.proj = nn.Linear(dim, dim)
60
+ self.proj_drop = nn.Dropout(proj_dropout)
61
+
62
+ def forward(
63
+ self,
64
+ x: torch.Tensor,
65
+ key_padding_mask: torch.Tensor | None = None,
66
+ ) -> torch.Tensor:
67
+ """Global self-attention on ``(B, N, C)`` tokens.
68
+
69
+ ``key_padding_mask`` (optional, shape ``(B, N)``, dtype bool):
70
+ ``True`` = real token, ``False`` = padding to ignore.
71
+ """
72
+ B, N, C = x.shape
73
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
74
+ qkv = qkv.permute(2, 0, 3, 1, 4)
75
+ q, k, v = qkv.unbind(dim=0)
76
+ q = self.q_norm(q)
77
+ k = self.k_norm(k)
78
+ attn_mask: torch.Tensor | None = None
79
+ if key_padding_mask is not None:
80
+ attn_mask = torch.zeros((B, 1, 1, N), dtype=q.dtype, device=q.device)
81
+ attn_mask = attn_mask.masked_fill(
82
+ ~key_padding_mask[:, None, None, :], float("-inf")
83
+ )
84
+ dropout_p = self.attention_dropout if self.training else 0.0
85
+ out = torch.nn.functional.scaled_dot_product_attention(
86
+ q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=False
87
+ )
88
+ out = out.transpose(1, 2).reshape(B, N, C)
89
+ return self.proj_drop(self.proj(out))
90
+
91
+
92
+ class TransformerBlock(nn.Module):
93
+ """Pre-norm transformer block with LayerScale and stochastic depth.
94
+
95
+ Residual pattern::
96
+
97
+ x = x + drop_path(γ_1 * Attn(LN(x)))
98
+ x = x + drop_path(γ_2 * MLP(LN(x)))
99
+
100
+ ``γ_*`` are per-channel learnable scales initialised near zero so
101
+ every block starts as an identity map.
102
+
103
+ Parameters
104
+ ----------
105
+ dim : int
106
+ Token dimension.
107
+ num_heads : int
108
+ Attention heads.
109
+ mlp_ratio : int, default=4
110
+ MLP hidden-dim multiplier.
111
+ dropout : float, default=0.0
112
+ MLP dropout.
113
+ attention_dropout : float, default=0.0
114
+ Attention-matrix dropout.
115
+ drop_path : float, default=0.0
116
+ Stochastic-depth probability for this block's residuals.
117
+ layerscale_init : float, default=1e-4
118
+ Initial value of the LayerScale gammas. Set to ``None`` to
119
+ disable LayerScale entirely.
120
+ """
121
+
122
+ def __init__(
123
+ self,
124
+ dim: int,
125
+ num_heads: int,
126
+ mlp_ratio: int = 4,
127
+ dropout: float = 0.0,
128
+ attention_dropout: float = 0.0,
129
+ drop_path: float = 0.0,
130
+ layerscale_init: float | None = 1e-4,
131
+ ) -> None:
132
+ super().__init__()
133
+ self.norm1 = nn.LayerNorm(dim)
134
+ self.attn = MultiHeadSelfAttention(
135
+ dim, num_heads, attention_dropout=attention_dropout, proj_dropout=dropout
136
+ )
137
+ self.drop_path1 = DropPath(drop_path)
138
+
139
+ self.norm2 = nn.LayerNorm(dim)
140
+ hidden = int(mlp_ratio * dim)
141
+ self.mlp = nn.Sequential(
142
+ nn.Linear(dim, hidden),
143
+ nn.GELU(),
144
+ nn.Dropout(dropout),
145
+ nn.Linear(hidden, dim),
146
+ nn.Dropout(dropout),
147
+ )
148
+ self.drop_path2 = DropPath(drop_path)
149
+
150
+ self.use_layerscale = layerscale_init is not None
151
+ if self.use_layerscale:
152
+ self.gamma1 = nn.Parameter(torch.full((dim,), float(layerscale_init)))
153
+ self.gamma2 = nn.Parameter(torch.full((dim,), float(layerscale_init)))
154
+
155
+ def forward(
156
+ self,
157
+ x: torch.Tensor,
158
+ key_padding_mask: torch.Tensor | None = None,
159
+ ) -> torch.Tensor:
160
+ """Run pre-norm attention + MLP residual sub-blocks with optional LayerScale."""
161
+ attn_out = self.attn(self.norm1(x), key_padding_mask=key_padding_mask)
162
+ mlp_out_src = self.norm2(x)
163
+ if self.use_layerscale:
164
+ attn_out = attn_out * self.gamma1
165
+ x = x + self.drop_path1(attn_out)
166
+ mlp_out = self.mlp(mlp_out_src)
167
+ if self.use_layerscale:
168
+ mlp_out = mlp_out * self.gamma2
169
+ x = x + self.drop_path2(mlp_out)
170
+ return x
171
+
172
+
173
+ class SpectralTransformer1D(nn.Module):
174
+ """1-D Vision Transformer backbone for binned spectra.
175
+
176
+ Parameters
177
+ ----------
178
+ input_dim : int
179
+ Number of input bins.
180
+ n_classes : int, default=2
181
+ Number of output logits.
182
+ patch_size : int, default=4
183
+ Non-overlapping patch width. Token count is
184
+ ``ceil(input_dim / patch_size)``.
185
+ embed_dim : int, default=64
186
+ Token embedding dimension.
187
+ depth : int, default=6
188
+ Number of transformer blocks.
189
+ num_heads : int, default=4
190
+ Attention heads per block. ``embed_dim`` must be divisible by
191
+ ``num_heads``.
192
+ mlp_ratio : int, default=4
193
+ MLP hidden-dim multiplier.
194
+ dropout : float, default=0.1
195
+ MLP dropout applied inside every block and before the head.
196
+ attention_dropout : float, default=0.0
197
+ Attention-matrix dropout.
198
+ drop_path_rate : float, default=0.1
199
+ End-of-stack stochastic-depth rate. Linearly interpolated
200
+ from ``0`` at block 0 to ``drop_path_rate`` at the final block.
201
+ layerscale_init : float or None, default=1e-4
202
+ LayerScale initial value. ``None`` disables LayerScale.
203
+ pool : {"cls", "mean"}, default="mean"
204
+ Aggregation strategy for classification. ``"mean"`` averages
205
+ over patch tokens (more robust on small data); ``"cls"``
206
+ prepends a learned token and uses its output.
207
+ head_dim : int, default=128
208
+ Width of the hidden dense layer in the classification head.
209
+ """
210
+
211
+ def __init__(
212
+ self,
213
+ input_dim: int,
214
+ n_classes: int = 2,
215
+ patch_size: int = 4,
216
+ embed_dim: int = 64,
217
+ depth: int = 6,
218
+ num_heads: int = 4,
219
+ mlp_ratio: int = 4,
220
+ dropout: float = 0.1,
221
+ attention_dropout: float = 0.0,
222
+ drop_path_rate: float = 0.1,
223
+ layerscale_init: float | None = 1e-4,
224
+ pool: str = "mean",
225
+ head_dim: int = 128,
226
+ ) -> None:
227
+ super().__init__()
228
+ if pool not in {"mean", "cls"}:
229
+ raise ValueError(f"pool must be 'mean' or 'cls'; got {pool!r}.")
230
+ if embed_dim % num_heads != 0:
231
+ raise ValueError(
232
+ f"embed_dim={embed_dim} must be divisible by num_heads={num_heads}."
233
+ )
234
+ if depth < 1:
235
+ raise ValueError(f"depth must be >= 1; got {depth!r}.")
236
+
237
+ self.pool = pool
238
+ self.patch_size = patch_size
239
+
240
+ self.embed = PatchEmbed1D(
241
+ patch_size=patch_size, in_channels=1, embed_dim=embed_dim
242
+ )
243
+
244
+ n_tokens = -(-input_dim // patch_size)
245
+ self.n_tokens = n_tokens
246
+
247
+ self.cls_token: nn.Parameter | None
248
+ if pool == "cls":
249
+ self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
250
+ nn.init.trunc_normal_(self.cls_token, std=0.02)
251
+ pos_len = n_tokens + 1
252
+ else:
253
+ self.register_parameter("cls_token", None)
254
+ pos_len = n_tokens
255
+ self.pos_embed = nn.Parameter(torch.zeros(1, pos_len, embed_dim))
256
+ nn.init.trunc_normal_(self.pos_embed, std=0.02)
257
+ self.pos_drop = nn.Dropout(dropout)
258
+
259
+ dpr = [float(x) for x in torch.linspace(0, drop_path_rate, depth)]
260
+ self.blocks = nn.ModuleList(
261
+ [
262
+ TransformerBlock(
263
+ dim=embed_dim,
264
+ num_heads=num_heads,
265
+ mlp_ratio=mlp_ratio,
266
+ dropout=dropout,
267
+ attention_dropout=attention_dropout,
268
+ drop_path=dpr[i],
269
+ layerscale_init=layerscale_init,
270
+ )
271
+ for i in range(depth)
272
+ ]
273
+ )
274
+ self.norm = nn.LayerNorm(embed_dim)
275
+ self.head = nn.Sequential(
276
+ nn.Linear(embed_dim, head_dim),
277
+ nn.GELU(),
278
+ nn.Dropout(dropout),
279
+ nn.Linear(head_dim, n_classes),
280
+ )
281
+
282
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
283
+ """Map ``(batch, input_dim)`` to ``(batch, n_classes)`` logits.
284
+
285
+ Notes
286
+ -----
287
+ Inputs whose length is not a multiple of ``patch_size`` are
288
+ right-padded with zeros before the patch embedding.
289
+ """
290
+ x = x.unsqueeze(1)
291
+ pad = (-x.shape[-1]) % self.patch_size
292
+ if pad:
293
+ x = torch.nn.functional.pad(x, (0, pad))
294
+ tokens = self.embed(x)
295
+ if self.cls_token is not None:
296
+ cls = self.cls_token.expand(tokens.shape[0], -1, -1)
297
+ tokens = torch.cat([cls, tokens], dim=1)
298
+ tokens = self.pos_drop(tokens + self.pos_embed[:, : tokens.shape[1]])
299
+ for block in self.blocks:
300
+ tokens = block(tokens)
301
+ tokens = self.norm(tokens)
302
+ if self.pool == "cls":
303
+ pooled = tokens[:, 0]
304
+ else:
305
+ start = 1 if self.cls_token is not None else 0
306
+ pooled = tokens[:, start:].mean(dim=1)
307
+ return self.head(pooled)
308
+
309
+
310
+ class MaldiTransformerClassifier(BaseSpectralClassifier):
311
+ """sklearn-compatible 1-D ViT classifier for MALDI-TOF spectra.
312
+
313
+ Parameters
314
+ ----------
315
+ patch_size : int, default=4
316
+ Patch size of the initial Conv1D embedding. Token count is
317
+ ``ceil(input_dim / patch_size)``.
318
+ embed_dim : int, default=64
319
+ Token embedding dimension. Must be divisible by ``num_heads``.
320
+ depth : int, default=6
321
+ Number of transformer blocks.
322
+ num_heads : int, default=4
323
+ Attention heads per block.
324
+ mlp_ratio : int, default=4
325
+ MLP hidden-dim multiplier inside each block.
326
+ dropout : float, default=0.1
327
+ MLP dropout applied inside every block and before the head.
328
+ attention_dropout : float, default=0.0
329
+ Attention-matrix dropout.
330
+ drop_path_rate : float, default=0.1
331
+ Linearly ramped stochastic-depth rate (0 at block 0, this
332
+ value at the final block).
333
+ layerscale_init : float or None, default=1e-4
334
+ LayerScale initial value. ``None`` disables LayerScale.
335
+ pool : {"mean", "cls"}, default="mean"
336
+ Token aggregation for classification.
337
+ head_dim : int, default=128
338
+ Width of the hidden dense layer in the classification head.
339
+ **kwargs
340
+ Forwarded to :class:`~maldideepkit.base.classifier.BaseSpectralClassifier`.
341
+
342
+ Notes
343
+ -----
344
+ Transformer training recipe baked in as defaults: ``lr=3e-4``,
345
+ ``weight_decay=0.05``, ``grad_clip_norm=1.0``, ``warmup_epochs=5``.
346
+
347
+ Examples
348
+ --------
349
+ >>> import numpy as np
350
+ >>> from maldideepkit import MaldiTransformerClassifier
351
+ >>> rng = np.random.default_rng(0)
352
+ >>> X = rng.standard_normal((32, 256)).astype("float32")
353
+ >>> y = rng.integers(0, 2, size=32)
354
+ >>> clf = MaldiTransformerClassifier(
355
+ ... epochs=2, batch_size=8, embed_dim=32, depth=2,
356
+ ... num_heads=2, patch_size=2, random_state=0,
357
+ ... ).fit(X, y)
358
+ >>> clf.predict(X).shape
359
+ (32,)
360
+ """
361
+
362
+ def __init__(
363
+ self,
364
+ input_dim: int | None = None,
365
+ n_classes: int = 2,
366
+ patch_size: int = 4,
367
+ embed_dim: int = 64,
368
+ depth: int = 6,
369
+ num_heads: int = 4,
370
+ mlp_ratio: int = 4,
371
+ dropout: float = 0.1,
372
+ attention_dropout: float = 0.0,
373
+ drop_path_rate: float = 0.1,
374
+ layerscale_init: float | None = 1e-4,
375
+ pool: str = "mean",
376
+ head_dim: int = 128,
377
+ learning_rate: float = 3e-4,
378
+ weight_decay: float = 0.05,
379
+ grad_clip_norm: float | None = 1.0,
380
+ label_smoothing: float = 0.0,
381
+ loss: str = "cross_entropy",
382
+ focal_gamma: float = 2.0,
383
+ use_amp: bool = False,
384
+ swa_start_epoch: int | None = None,
385
+ tune_threshold: bool = False,
386
+ threshold_metric: str = "balanced_accuracy",
387
+ calibrate_temperature: bool = False,
388
+ min_val_auroc_for_threshold_tune: float = 0.6,
389
+ use_sam: bool = False,
390
+ sam_rho: float = 0.05,
391
+ batch_size: int = 32,
392
+ epochs: int = 100,
393
+ early_stopping_patience: int = 10,
394
+ val_fraction: float = 0.1,
395
+ warmup_epochs: int = 5,
396
+ standardize: bool = False,
397
+ input_transform: str | None = None,
398
+ warping: Any | None = None,
399
+ metrics_log_path: str | Path | None = None,
400
+ track_train_metrics: bool = False,
401
+ augment: Any | None = None,
402
+ mixup_alpha: float = 0.0,
403
+ cutmix_alpha: float = 0.0,
404
+ ema_decay: float | None = None,
405
+ retry_on_val_auroc_below: float | None = None,
406
+ max_retries: int = 2,
407
+ class_weight: str | np.ndarray | list | None = None,
408
+ device: str | torch.device = "auto",
409
+ random_state: int = 0,
410
+ verbose: bool = False,
411
+ ) -> None:
412
+ super().__init__(
413
+ input_dim=input_dim,
414
+ n_classes=n_classes,
415
+ learning_rate=learning_rate,
416
+ weight_decay=weight_decay,
417
+ grad_clip_norm=grad_clip_norm,
418
+ label_smoothing=label_smoothing,
419
+ loss=loss,
420
+ focal_gamma=focal_gamma,
421
+ use_amp=use_amp,
422
+ swa_start_epoch=swa_start_epoch,
423
+ tune_threshold=tune_threshold,
424
+ threshold_metric=threshold_metric,
425
+ calibrate_temperature=calibrate_temperature,
426
+ min_val_auroc_for_threshold_tune=min_val_auroc_for_threshold_tune,
427
+ use_sam=use_sam,
428
+ sam_rho=sam_rho,
429
+ batch_size=batch_size,
430
+ epochs=epochs,
431
+ early_stopping_patience=early_stopping_patience,
432
+ val_fraction=val_fraction,
433
+ warmup_epochs=warmup_epochs,
434
+ standardize=standardize,
435
+ input_transform=input_transform,
436
+ warping=warping,
437
+ metrics_log_path=metrics_log_path,
438
+ track_train_metrics=track_train_metrics,
439
+ augment=augment,
440
+ mixup_alpha=mixup_alpha,
441
+ cutmix_alpha=cutmix_alpha,
442
+ ema_decay=ema_decay,
443
+ retry_on_val_auroc_below=retry_on_val_auroc_below,
444
+ max_retries=max_retries,
445
+ class_weight=class_weight,
446
+ device=device,
447
+ random_state=random_state,
448
+ verbose=verbose,
449
+ )
450
+ self.patch_size = patch_size
451
+ self.embed_dim = embed_dim
452
+ self.depth = depth
453
+ self.num_heads = num_heads
454
+ self.mlp_ratio = mlp_ratio
455
+ self.dropout = dropout
456
+ self.attention_dropout = attention_dropout
457
+ self.drop_path_rate = drop_path_rate
458
+ self.layerscale_init = layerscale_init
459
+ self.pool = pool
460
+ self.head_dim = head_dim
461
+
462
+ def _build_model(self) -> nn.Module:
463
+ return SpectralTransformer1D(
464
+ input_dim=self.input_dim_,
465
+ n_classes=self.n_classes_,
466
+ patch_size=int(self.patch_size),
467
+ embed_dim=int(self.embed_dim),
468
+ depth=int(self.depth),
469
+ num_heads=int(self.num_heads),
470
+ mlp_ratio=int(self.mlp_ratio),
471
+ dropout=float(self.dropout),
472
+ attention_dropout=float(self.attention_dropout),
473
+ drop_path_rate=float(self.drop_path_rate),
474
+ layerscale_init=self.layerscale_init,
475
+ pool=str(self.pool),
476
+ head_dim=int(self.head_dim),
477
+ )
478
+
479
+ @classmethod
480
+ def from_spectrum(
481
+ cls, bin_width: int, input_dim: int, **overrides
482
+ ) -> "MaldiTransformerClassifier":
483
+ """Construct a classifier for a given ``(bin_width, input_dim)`` layout.
484
+
485
+ The transformer is architecturally scale-agnostic, so this
486
+ factory only forwards ``input_dim`` and any ``**overrides``.
487
+ Provided for API symmetry with the other classifiers.
488
+ """
489
+ del bin_width
490
+ kwargs: dict[str, Any] = {"input_dim": input_dim}
491
+ kwargs.update(overrides)
492
+ return cls(**kwargs)
@@ -0,0 +1,22 @@
1
+ """Reproducibility and training helpers shared across model families."""
2
+
3
+ from .calibration import fit_temperature, tune_threshold
4
+ from .ensemble import SpectralEnsemble
5
+ from .loss import FocalLoss
6
+ from .lr_finder import find_lr
7
+ from .reproducibility import resolve_device, seed_everything
8
+ from .sam import SAMOptimizer
9
+ from .training import EarlyStopping, train_loop
10
+
11
+ __all__ = [
12
+ "EarlyStopping",
13
+ "FocalLoss",
14
+ "SAMOptimizer",
15
+ "SpectralEnsemble",
16
+ "find_lr",
17
+ "fit_temperature",
18
+ "resolve_device",
19
+ "seed_everything",
20
+ "train_loop",
21
+ "tune_threshold",
22
+ ]
@@ -0,0 +1,134 @@
1
+ """Post-hoc calibration helpers used by :class:`BaseSpectralClassifier`.
2
+
3
+ - :func:`tune_threshold` picks the binary decision threshold on a
4
+ validation set that maximises a chosen metric.
5
+ - :func:`fit_temperature` optimises a single temperature scalar by
6
+ LBFGS on held-out logits for probability calibration.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import numpy as np
12
+ import torch
13
+ import torch.nn.functional as F
14
+ from sklearn.metrics import balanced_accuracy_score, f1_score, roc_curve
15
+
16
+
17
+ def tune_threshold(
18
+ y_true: np.ndarray,
19
+ y_proba: np.ndarray,
20
+ metric: str = "balanced_accuracy",
21
+ ) -> float:
22
+ """Pick the binary decision threshold that maximises ``metric``.
23
+
24
+ Sweeps the unique observed probabilities (capped at 1000 quantiles)
25
+ so severely-imbalanced settings still resolve. Falls back to a
26
+ 99-point ``linspace(0.01, 0.99)`` only when no probability lies
27
+ strictly inside ``(0, 1)``.
28
+
29
+ Parameters
30
+ ----------
31
+ y_true : array-like of shape (n_samples,)
32
+ Binary ground-truth labels in ``{0, 1}``.
33
+ y_proba : array-like of shape (n_samples,) or (n_samples, 2)
34
+ Predicted positive-class probabilities. If a 2-D array is
35
+ given, column index ``1`` is used.
36
+ metric : {"balanced_accuracy", "f1", "youden"}, default="balanced_accuracy"
37
+ Which metric to maximise. ``"youden"`` = TPR - FPR.
38
+
39
+ Returns
40
+ -------
41
+ float
42
+ Threshold in ``(0, 1)``. Use as ``y_pred = (y_proba >= t)``.
43
+ """
44
+ y_true = np.asarray(y_true).ravel().astype(int)
45
+ y_proba_arr = np.asarray(y_proba, dtype=float)
46
+ if y_proba_arr.ndim == 2:
47
+ if y_proba_arr.shape[1] != 2:
48
+ raise ValueError(
49
+ "tune_threshold is binary-only; "
50
+ f"got y_proba with {y_proba_arr.shape[1]} columns."
51
+ )
52
+ y_proba_arr = y_proba_arr[:, 1]
53
+ y_proba_arr = y_proba_arr.ravel()
54
+
55
+ if metric == "youden":
56
+ fpr, tpr, thr = roc_curve(y_true, y_proba_arr)
57
+ valid = (thr > 0) & (thr < 1)
58
+ if not valid.any():
59
+ return 0.5
60
+ j = tpr[valid] - fpr[valid]
61
+ return float(thr[valid][int(np.argmax(j))])
62
+
63
+ unique = np.unique(y_proba_arr)
64
+ unique = unique[(unique > 0) & (unique < 1)]
65
+ if unique.size == 0:
66
+ candidates = np.linspace(0.01, 0.99, 99)
67
+ elif unique.size > 1000:
68
+ candidates = np.quantile(unique, np.linspace(0.0, 1.0, 1000))
69
+ else:
70
+ candidates = unique
71
+ best_t, best_score = 0.5, -np.inf
72
+ for t in candidates:
73
+ pred = (y_proba_arr >= t).astype(int)
74
+ if metric == "balanced_accuracy":
75
+ score = balanced_accuracy_score(y_true, pred)
76
+ elif metric == "f1":
77
+ score = f1_score(y_true, pred, zero_division=0)
78
+ else:
79
+ raise ValueError(
80
+ f"Unknown metric={metric!r}; "
81
+ "expected 'balanced_accuracy', 'f1', or 'youden'."
82
+ )
83
+ if score > best_score:
84
+ best_score, best_t = score, float(t)
85
+ return best_t
86
+
87
+
88
+ def fit_temperature(
89
+ logits: torch.Tensor | np.ndarray,
90
+ y_true: torch.Tensor | np.ndarray,
91
+ max_iter: int = 200,
92
+ lr: float = 1e-1,
93
+ ) -> float:
94
+ """Fit a scalar temperature by LBFGS minimisation of NLL.
95
+
96
+ Applies to raw logits (not probabilities). Returns the temperature
97
+ ``T`` such that ``softmax(logits / T)`` is better-calibrated than
98
+ the unscaled softmax.
99
+
100
+ Parameters
101
+ ----------
102
+ logits : torch.Tensor or ndarray of shape (n_samples, n_classes)
103
+ Held-out logits.
104
+ y_true : torch.Tensor or ndarray of shape (n_samples,)
105
+ Ground-truth class indices.
106
+ max_iter : int, default=200
107
+ LBFGS max iterations.
108
+ lr : float, default=1e-1
109
+ LBFGS step size.
110
+
111
+ Returns
112
+ -------
113
+ float
114
+ Fitted temperature; strictly positive.
115
+ """
116
+ if not isinstance(logits, torch.Tensor):
117
+ logits = torch.as_tensor(logits, dtype=torch.float32)
118
+ if not isinstance(y_true, torch.Tensor):
119
+ y_true = torch.as_tensor(np.asarray(y_true).ravel(), dtype=torch.long)
120
+ else:
121
+ y_true = y_true.to(torch.long).view(-1)
122
+
123
+ log_temperature = torch.zeros(1, device=logits.device, requires_grad=True)
124
+ optimizer = torch.optim.LBFGS([log_temperature], lr=lr, max_iter=max_iter)
125
+
126
+ def _closure():
127
+ optimizer.zero_grad()
128
+ t = torch.exp(log_temperature)
129
+ loss = F.cross_entropy(logits / t, y_true)
130
+ loss.backward()
131
+ return loss
132
+
133
+ optimizer.step(_closure)
134
+ return float(torch.exp(log_temperature).detach().item())