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.
- maldideepkit/__init__.py +54 -0
- maldideepkit/_bin_scaling.py +52 -0
- maldideepkit/_blocks.py +80 -0
- maldideepkit/attention/__init__.py +5 -0
- maldideepkit/attention/mlp.py +319 -0
- maldideepkit/augment/__init__.py +15 -0
- maldideepkit/augment/mixing.py +113 -0
- maldideepkit/augment/spectra.py +251 -0
- maldideepkit/base/__init__.py +15 -0
- maldideepkit/base/classifier.py +1079 -0
- maldideepkit/base/data.py +322 -0
- maldideepkit/blocks.py +39 -0
- maldideepkit/cnn/__init__.py +5 -0
- maldideepkit/cnn/cnn.py +316 -0
- maldideepkit/py.typed +0 -0
- maldideepkit/resnet/__init__.py +5 -0
- maldideepkit/resnet/resnet.py +380 -0
- maldideepkit/transformer/__init__.py +7 -0
- maldideepkit/transformer/transformer.py +492 -0
- maldideepkit/utils/__init__.py +22 -0
- maldideepkit/utils/calibration.py +134 -0
- maldideepkit/utils/ensemble.py +132 -0
- maldideepkit/utils/loss.py +138 -0
- maldideepkit/utils/lr_finder.py +173 -0
- maldideepkit/utils/reproducibility.py +70 -0
- maldideepkit/utils/sam.py +121 -0
- maldideepkit/utils/training.py +386 -0
- maldideepkit-0.1.0.dist-info/METADATA +301 -0
- maldideepkit-0.1.0.dist-info/RECORD +32 -0
- maldideepkit-0.1.0.dist-info/WHEEL +5 -0
- maldideepkit-0.1.0.dist-info/licenses/LICENSE +21 -0
- maldideepkit-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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())
|