islkit 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.
- islkit/__init__.py +88 -0
- islkit/adapters.py +237 -0
- islkit/baseline.py +287 -0
- islkit/data.py +320 -0
- islkit/device.py +38 -0
- islkit/domain.py +187 -0
- islkit/features.py +323 -0
- islkit/infer.py +911 -0
- islkit/labels.py +213 -0
- islkit/metrics.py +81 -0
- islkit/model.py +623 -0
- islkit/pipeline.py +717 -0
- islkit/plotting.py +131 -0
- islkit/seeding.py +19 -0
- islkit/server.py +246 -0
- islkit/view.py +287 -0
- islkit/viz.py +435 -0
- islkit-0.1.0.dist-info/METADATA +200 -0
- islkit-0.1.0.dist-info/RECORD +21 -0
- islkit-0.1.0.dist-info/WHEEL +4 -0
- islkit-0.1.0.dist-info/licenses/LICENSE +21 -0
islkit/model.py
ADDED
|
@@ -0,0 +1,623 @@
|
|
|
1
|
+
"""The dual-branch TCN: the pretraining model, and the backbone fine-tuning reuses.
|
|
2
|
+
|
|
3
|
+
**Why two stacks and not one.** Handshape and location have different temporal
|
|
4
|
+
dynamics. A handshape is close to a held configuration — it changes at
|
|
5
|
+
articulation boundaries and is otherwise flat, so what a filter needs to find is
|
|
6
|
+
*which* configuration is being held. A wrist trajectory is the opposite: it is
|
|
7
|
+
motion throughout, and what matters is its shape over tens of frames. One shared
|
|
8
|
+
filter bank has to compromise between those, and the compromise costs both. Two
|
|
9
|
+
stacks with their own dilations do not.
|
|
10
|
+
|
|
11
|
+
The split of the 352-dim frame vector:
|
|
12
|
+
|
|
13
|
+
branch A hand-local shapes 126 + velocity 126 = 252 channels
|
|
14
|
+
branch B wrist positions + pose 39 + velocity 39 = 78 channels
|
|
15
|
+
head non-manual scalars 4 + velocity 4 = 8, pooled straight
|
|
16
|
+
(gates) validity mask 14 multiplied through
|
|
17
|
+
|
|
18
|
+
Those are all 352. The mask is the part that is easy to miss: it is not a
|
|
19
|
+
channel in either branch. A validity bit is carried per part and **multiplied
|
|
20
|
+
through**, because when the mask was offered to the XGBoost baseline as a feature
|
|
21
|
+
the trees ignored it (0.1% of gain). Here each geometry dim is multiplied by the bit for the part it
|
|
22
|
+
came from, so an absent hand contributes exactly zero rather than a plausible
|
|
23
|
+
hand-shaped vector left over from resampling.
|
|
24
|
+
|
|
25
|
+
That matters more than it looks. `encode_clip` linearly resamples all 183 dims to
|
|
26
|
+
T=48, the mask included, so on real clips a hand appearing mid-clip leaves
|
|
27
|
+
*fractional* mask values and interpolated geometry at the transition (on INCLUDE,
|
|
28
|
+
8.9% of hand-presence bits are neither 0 nor 1). Gating re-attenuates that
|
|
29
|
+
geometry by exactly the confidence the mask carries.
|
|
30
|
+
|
|
31
|
+
**Normalise first, then gate.** Each source dim goes through a BatchNorm before
|
|
32
|
+
the gate is applied, so "absent" lands on 0 — the neutral value — rather than on
|
|
33
|
+
whatever -mean/std happens to be. The BatchNorm is also the honest answer to the
|
|
34
|
+
question of z's scale: after body-frame normalisation z is the
|
|
35
|
+
highest-variance channel in the pose block (2.8x x), and per-channel
|
|
36
|
+
standardisation removes that spread before any filter sees it. A tree splits on
|
|
37
|
+
separation and so ignored it; a convolution weights by magnitude and would not.
|
|
38
|
+
|
|
39
|
+
**Backbone + swappable head, from the start.** Fine-tuning replaces the head with
|
|
40
|
+
a ~40-class one, so `replace_head` and `freeze_backbone` exist from the start
|
|
41
|
+
rather than being retrofitted around whatever pretraining happened to build. The dense
|
|
42
|
+
projection is part of the *backbone*, not the head: it is task-agnostic, and
|
|
43
|
+
keeping it on the frozen side is what makes frozen-backbone-plus-head a real
|
|
44
|
+
option. Transfer measured on synthetic domains says that option loses nothing.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
from __future__ import annotations
|
|
48
|
+
|
|
49
|
+
import math
|
|
50
|
+
import time
|
|
51
|
+
from dataclasses import dataclass, field
|
|
52
|
+
from pathlib import Path
|
|
53
|
+
|
|
54
|
+
import numpy as np
|
|
55
|
+
import torch
|
|
56
|
+
import torch.nn as nn
|
|
57
|
+
|
|
58
|
+
from islkit.features import DIM_FRAME, DIM_GEOM, N_HAND, N_POSE_SEL
|
|
59
|
+
|
|
60
|
+
# --------------------------------------------------------------------------
|
|
61
|
+
# Where each dim of the 352-vector goes, and which mask bit gates it
|
|
62
|
+
# --------------------------------------------------------------------------
|
|
63
|
+
|
|
64
|
+
# The mask block, [169:183], in the order encode_frame writes it.
|
|
65
|
+
MASK_START, MASK_END = DIM_GEOM, DIM_FRAME # 169, 183
|
|
66
|
+
GATE_HAND = (0, 1) # slot 0 / slot 1 present
|
|
67
|
+
GATE_POSE_0 = 2 # first of N_POSE_SEL visibility bits
|
|
68
|
+
GATE_FACE = MASK_END - MASK_START - 1 # 13, the face-present bit
|
|
69
|
+
|
|
70
|
+
_HAND = N_HAND * 3 # 63, one hand's local shape
|
|
71
|
+
_POSE = N_POSE_SEL * 3 # 33, upper-body pose
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _branch_indices() -> dict[str, tuple[list[int], list[int]]]:
|
|
75
|
+
"""Source dims and their gating mask bit, per branch.
|
|
76
|
+
|
|
77
|
+
Returned rather than hard-coded so the layout is derived from features.py's
|
|
78
|
+
own constants. A wrong offset here routes the face mesh into the handshape
|
|
79
|
+
branch and still trains, still converges, and still prints a number — the
|
|
80
|
+
same failure mode the parquet block-routing test exists for.
|
|
81
|
+
"""
|
|
82
|
+
a_src, a_gate = [], []
|
|
83
|
+
for base in (0, DIM_FRAME): # hand_shape at 0, d_hand_shape at 183
|
|
84
|
+
for slot in (0, 1):
|
|
85
|
+
a_src.extend(range(base + slot * _HAND, base + (slot + 1) * _HAND))
|
|
86
|
+
a_gate.extend([GATE_HAND[slot]] * _HAND)
|
|
87
|
+
|
|
88
|
+
b_src, b_gate = [], []
|
|
89
|
+
wrist_0 = 2 * _HAND # 126
|
|
90
|
+
pose_0 = wrist_0 + 6 # 132
|
|
91
|
+
for base in (0, DIM_FRAME):
|
|
92
|
+
for slot in (0, 1): # wrist position shares the hand's validity bit
|
|
93
|
+
b_src.extend(range(base + wrist_0 + slot * 3, base + wrist_0 + (slot + 1) * 3))
|
|
94
|
+
b_gate.extend([GATE_HAND[slot]] * 3)
|
|
95
|
+
for lm in range(N_POSE_SEL): # each pose landmark has its own visibility
|
|
96
|
+
b_src.extend(range(base + pose_0 + lm * 3, base + pose_0 + (lm + 1) * 3))
|
|
97
|
+
b_gate.extend([GATE_POSE_0 + lm] * 3)
|
|
98
|
+
|
|
99
|
+
n_src, n_gate = [], []
|
|
100
|
+
nm_0 = pose_0 + _POSE # 165
|
|
101
|
+
for base in (0, DIM_FRAME):
|
|
102
|
+
n_src.extend(range(base + nm_0, base + nm_0 + 4))
|
|
103
|
+
n_gate.extend([GATE_FACE] * 4)
|
|
104
|
+
|
|
105
|
+
return {"shape": (a_src, a_gate), "motion": (b_src, b_gate), "non_manual": (n_src, n_gate)}
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
BRANCH_INDICES = _branch_indices()
|
|
109
|
+
DIM_SHAPE = len(BRANCH_INDICES["shape"][0]) # 252
|
|
110
|
+
DIM_MOTION = len(BRANCH_INDICES["motion"][0]) # 78
|
|
111
|
+
DIM_NON_MANUAL = len(BRANCH_INDICES["non_manual"][0]) # 8
|
|
112
|
+
DIM_INPUT = DIM_FRAME + DIM_GEOM # 352
|
|
113
|
+
|
|
114
|
+
DILATIONS: tuple[int, ...] = (1, 2, 4, 8, 16)
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
# --------------------------------------------------------------------------
|
|
118
|
+
# Building blocks
|
|
119
|
+
# --------------------------------------------------------------------------
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class TCNBlock(nn.Module):
|
|
123
|
+
"""Two dilated convolutions, then add the input back.
|
|
124
|
+
|
|
125
|
+
`padding = dilation * (kernel - 1) // 2` keeps the time axis the same length,
|
|
126
|
+
so blocks stack without any shape bookkeeping. Residual because with five
|
|
127
|
+
blocks the useful thing for a later block to learn is a correction, not a
|
|
128
|
+
replacement.
|
|
129
|
+
"""
|
|
130
|
+
|
|
131
|
+
def __init__(self, ch: int, dilation: int, kernel: int = 3, dropout: float = 0.15):
|
|
132
|
+
super().__init__()
|
|
133
|
+
pad = dilation * (kernel - 1) // 2
|
|
134
|
+
self.conv1 = nn.Conv1d(ch, ch, kernel, padding=pad, dilation=dilation)
|
|
135
|
+
self.conv2 = nn.Conv1d(ch, ch, kernel, padding=pad, dilation=dilation)
|
|
136
|
+
self.norm1 = nn.BatchNorm1d(ch)
|
|
137
|
+
self.norm2 = nn.BatchNorm1d(ch)
|
|
138
|
+
self.act = nn.ReLU()
|
|
139
|
+
self.drop = nn.Dropout(dropout)
|
|
140
|
+
|
|
141
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
142
|
+
h = self.drop(self.act(self.norm1(self.conv1(x))))
|
|
143
|
+
h = self.drop(self.act(self.norm2(self.conv2(h))))
|
|
144
|
+
return self.act(x + h)
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
class DilatedStack(nn.Module):
|
|
148
|
+
"""One branch: 1x1 stem to `ch` channels, then the dilated residual blocks.
|
|
149
|
+
|
|
150
|
+
Input arrives as (B, T, C) — the natural shape out of `encode_clip` — and is
|
|
151
|
+
transposed to Conv1d's (B, C, T) here, once, in the one place that owns it.
|
|
152
|
+
"""
|
|
153
|
+
|
|
154
|
+
def __init__(
|
|
155
|
+
self,
|
|
156
|
+
c_in: int,
|
|
157
|
+
ch: int,
|
|
158
|
+
dilations: tuple[int, ...] = DILATIONS,
|
|
159
|
+
kernel: int = 3,
|
|
160
|
+
dropout: float = 0.15,
|
|
161
|
+
):
|
|
162
|
+
super().__init__()
|
|
163
|
+
self.stem = nn.Conv1d(c_in, ch, 1)
|
|
164
|
+
self.blocks = nn.Sequential(*[TCNBlock(ch, d, kernel, dropout) for d in dilations])
|
|
165
|
+
self.out_ch = ch
|
|
166
|
+
self.receptive_field = 1 + sum(2 * (kernel - 1) * d for d in dilations)
|
|
167
|
+
|
|
168
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
169
|
+
return self.blocks(self.stem(x.transpose(1, 2)))
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
# --------------------------------------------------------------------------
|
|
173
|
+
# The backbone
|
|
174
|
+
# --------------------------------------------------------------------------
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
class Backbone(nn.Module):
|
|
178
|
+
"""(B, T, 352) -> (B, embed). Everything fine-tuning keeps.
|
|
179
|
+
|
|
180
|
+
Two dilated stacks over the two branches, mean-pooled over time and
|
|
181
|
+
concatenated with the time-pooled non-manual scalars, then projected to a
|
|
182
|
+
fixed-width embedding. The classifier on top is one Linear, and is the only
|
|
183
|
+
part `replace_head` touches.
|
|
184
|
+
"""
|
|
185
|
+
|
|
186
|
+
def __init__(
|
|
187
|
+
self,
|
|
188
|
+
ch_shape: int = 96,
|
|
189
|
+
ch_motion: int = 64,
|
|
190
|
+
embed: int = 192,
|
|
191
|
+
dilations: tuple[int, ...] = DILATIONS,
|
|
192
|
+
kernel: int = 3,
|
|
193
|
+
dropout: float = 0.15,
|
|
194
|
+
mask_gate: bool = True,
|
|
195
|
+
):
|
|
196
|
+
super().__init__()
|
|
197
|
+
self.config = {
|
|
198
|
+
"ch_shape": ch_shape,
|
|
199
|
+
"ch_motion": ch_motion,
|
|
200
|
+
"embed": embed,
|
|
201
|
+
"dilations": list(dilations),
|
|
202
|
+
"kernel": kernel,
|
|
203
|
+
"dropout": dropout,
|
|
204
|
+
"mask_gate": mask_gate,
|
|
205
|
+
}
|
|
206
|
+
self.mask_gate = mask_gate
|
|
207
|
+
self._frozen = False
|
|
208
|
+
|
|
209
|
+
for name in ("shape", "motion", "non_manual"):
|
|
210
|
+
src, gate = BRANCH_INDICES[name]
|
|
211
|
+
self.register_buffer(f"{name}_src", torch.tensor(src, dtype=torch.long))
|
|
212
|
+
self.register_buffer(f"{name}_gate", torch.tensor(gate, dtype=torch.long))
|
|
213
|
+
|
|
214
|
+
# Per-source-dim standardisation, applied before the gate so that
|
|
215
|
+
# "absent" is 0 and not an arbitrary offset. See the module docstring.
|
|
216
|
+
self.norm_shape = nn.BatchNorm1d(DIM_SHAPE)
|
|
217
|
+
self.norm_motion = nn.BatchNorm1d(DIM_MOTION)
|
|
218
|
+
self.norm_non_manual = nn.BatchNorm1d(DIM_NON_MANUAL)
|
|
219
|
+
|
|
220
|
+
self.shape_stack = DilatedStack(DIM_SHAPE, ch_shape, dilations, kernel, dropout)
|
|
221
|
+
self.motion_stack = DilatedStack(DIM_MOTION, ch_motion, dilations, kernel, dropout)
|
|
222
|
+
|
|
223
|
+
self.project = nn.Sequential(
|
|
224
|
+
nn.Linear(ch_shape + ch_motion + DIM_NON_MANUAL, embed),
|
|
225
|
+
nn.BatchNorm1d(embed),
|
|
226
|
+
nn.ReLU(),
|
|
227
|
+
nn.Dropout(dropout),
|
|
228
|
+
)
|
|
229
|
+
self.out_dim = embed
|
|
230
|
+
self.receptive_field = self.shape_stack.receptive_field
|
|
231
|
+
|
|
232
|
+
def _branch(self, x: torch.Tensor, name: str, norm: nn.Module) -> torch.Tensor:
|
|
233
|
+
"""Gather one branch's dims, standardise them, then gate by the mask."""
|
|
234
|
+
src = getattr(self, f"{name}_src")
|
|
235
|
+
out = x.index_select(-1, src)
|
|
236
|
+
out = norm(out.transpose(1, 2)).transpose(1, 2) # BatchNorm1d wants (B, C, T)
|
|
237
|
+
if self.mask_gate:
|
|
238
|
+
gate = x[..., MASK_START:MASK_END].index_select(-1, getattr(self, f"{name}_gate"))
|
|
239
|
+
out = out * gate
|
|
240
|
+
return out
|
|
241
|
+
|
|
242
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
243
|
+
if x.shape[-1] != DIM_INPUT:
|
|
244
|
+
raise ValueError(f"expected (B, T, {DIM_INPUT}) features, got {tuple(x.shape)}")
|
|
245
|
+
|
|
246
|
+
a = self.shape_stack(self._branch(x, "shape", self.norm_shape)).mean(dim=2)
|
|
247
|
+
b = self.motion_stack(self._branch(x, "motion", self.norm_motion)).mean(dim=2)
|
|
248
|
+
# The four non-manual scalars never enter a convolution. Brow raise and
|
|
249
|
+
# mouth shape are held configurations, not trajectories, and the baseline measured
|
|
250
|
+
# them at 1.4% of gain — a dilated stack over four dims would be four
|
|
251
|
+
# times the parameters for something a mean already says.
|
|
252
|
+
n = self._branch(x, "non_manual", self.norm_non_manual).mean(dim=1)
|
|
253
|
+
|
|
254
|
+
return self.project(torch.cat([a, b, n], dim=1))
|
|
255
|
+
|
|
256
|
+
def train(self, mode: bool = True): # noqa: D102 — nn.Module override
|
|
257
|
+
# A frozen backbone must stay in eval: BatchNorm updates its running
|
|
258
|
+
# statistics from the forward pass regardless of requires_grad, so a
|
|
259
|
+
# "frozen" backbone left in train mode drifts on the fine-tuning set and
|
|
260
|
+
# the frozen-head result stops being reproducible.
|
|
261
|
+
return super().train(mode and not self._frozen)
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
class SignClassifier(nn.Module):
|
|
265
|
+
"""Backbone plus a linear head. The head is the part that gets replaced."""
|
|
266
|
+
|
|
267
|
+
def __init__(self, backbone: Backbone, n_classes: int):
|
|
268
|
+
super().__init__()
|
|
269
|
+
self.backbone = backbone
|
|
270
|
+
self.head = nn.Linear(backbone.out_dim, n_classes)
|
|
271
|
+
self.n_classes = n_classes
|
|
272
|
+
|
|
273
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
274
|
+
return self.head(self.backbone(x))
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def build_model(n_classes: int, **kwargs) -> SignClassifier:
|
|
278
|
+
"""A fresh classifier. `kwargs` go to `Backbone`."""
|
|
279
|
+
return SignClassifier(Backbone(**kwargs), n_classes)
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def replace_head(model: SignClassifier, n_classes: int) -> SignClassifier:
|
|
283
|
+
"""Swap in a fresh head of a different width, keeping the backbone.
|
|
284
|
+
|
|
285
|
+
Fine-tuning's first move: 262 pretraining classes out, ~40 device classes in. The new
|
|
286
|
+
head lands on the same device as the old one, which is the kind of detail
|
|
287
|
+
that otherwise surfaces as a device-mismatch error twenty minutes in.
|
|
288
|
+
"""
|
|
289
|
+
device = next(model.head.parameters()).device
|
|
290
|
+
model.head = nn.Linear(model.backbone.out_dim, n_classes).to(device)
|
|
291
|
+
model.n_classes = n_classes
|
|
292
|
+
return model
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def freeze_backbone(model: SignClassifier, freeze: bool = True) -> SignClassifier:
|
|
296
|
+
"""Freeze (or unfreeze) everything except the head.
|
|
297
|
+
|
|
298
|
+
Both `requires_grad` and BatchNorm mode, because the second is the one that
|
|
299
|
+
is easy to forget and silently makes the run unreproducible — see
|
|
300
|
+
`Backbone.train`.
|
|
301
|
+
"""
|
|
302
|
+
for p in model.backbone.parameters():
|
|
303
|
+
p.requires_grad = not freeze
|
|
304
|
+
model.backbone._frozen = freeze
|
|
305
|
+
model.backbone.train(model.training)
|
|
306
|
+
return model
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
def n_parameters(model: nn.Module, trainable_only: bool = False) -> int:
|
|
310
|
+
ps = model.parameters()
|
|
311
|
+
return sum(p.numel() for p in ps if p.requires_grad or not trainable_only)
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
# --------------------------------------------------------------------------
|
|
315
|
+
# Saving what fine-tuning needs
|
|
316
|
+
# --------------------------------------------------------------------------
|
|
317
|
+
|
|
318
|
+
|
|
319
|
+
def save_backbone(backbone: Backbone, path: str | Path, training: dict | None = None) -> Path:
|
|
320
|
+
"""Weights, the config that built them, and how they were trained.
|
|
321
|
+
|
|
322
|
+
The config travels with the weights for the same reason the label map does
|
|
323
|
+
(both are frozen at training time): reconstructing either from code that has moved on since
|
|
324
|
+
is how a checkpoint loads cleanly and means something else.
|
|
325
|
+
|
|
326
|
+
`training` records the run that produced these weights — learning rate,
|
|
327
|
+
epochs, seed, clip count. The architecture config alone cannot distinguish
|
|
328
|
+
two checkpoints trained with different hyperparameters, so a re-run silently
|
|
329
|
+
replaces the file that a committed summary.json claims to describe and
|
|
330
|
+
nothing anywhere disagrees. `data.py` fingerprints the feature cache against
|
|
331
|
+
exactly this failure; weights are worth the same treatment.
|
|
332
|
+
"""
|
|
333
|
+
path = Path(path)
|
|
334
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
335
|
+
torch.save(
|
|
336
|
+
{
|
|
337
|
+
"config": backbone.config,
|
|
338
|
+
"state_dict": backbone.state_dict(),
|
|
339
|
+
"training": dict(training or {}),
|
|
340
|
+
},
|
|
341
|
+
path,
|
|
342
|
+
)
|
|
343
|
+
return path
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
def describe_backbone(path: str | Path) -> dict:
|
|
347
|
+
"""What produced a checkpoint, without building the model.
|
|
348
|
+
|
|
349
|
+
For answering "is the file on disk the one the write-up describes?" — which
|
|
350
|
+
otherwise needs a timestamp and a good memory.
|
|
351
|
+
"""
|
|
352
|
+
blob = torch.load(Path(path), map_location="cpu", weights_only=True)
|
|
353
|
+
return {"config": blob["config"], "training": blob.get("training", {})}
|
|
354
|
+
|
|
355
|
+
|
|
356
|
+
def load_backbone(path: str | Path, device=None) -> Backbone:
|
|
357
|
+
"""Rebuild a backbone from `save_backbone` output."""
|
|
358
|
+
blob = torch.load(Path(path), map_location=device or "cpu", weights_only=True)
|
|
359
|
+
cfg = dict(blob["config"])
|
|
360
|
+
cfg["dilations"] = tuple(cfg["dilations"])
|
|
361
|
+
backbone = Backbone(**cfg)
|
|
362
|
+
backbone.load_state_dict(blob["state_dict"])
|
|
363
|
+
return backbone.to(device) if device is not None else backbone
|
|
364
|
+
|
|
365
|
+
|
|
366
|
+
# Sentinel distinguishing "not passed" (compute it) from an explicit `None`
|
|
367
|
+
# ("write a checkpoint without one" — the pre-guard shape `load_classifier`
|
|
368
|
+
# callers must still be able to load).
|
|
369
|
+
_COMPUTE = object()
|
|
370
|
+
|
|
371
|
+
|
|
372
|
+
def save_classifier(
|
|
373
|
+
model: SignClassifier,
|
|
374
|
+
path: str | Path,
|
|
375
|
+
*,
|
|
376
|
+
labels: str,
|
|
377
|
+
encoder: dict,
|
|
378
|
+
training: dict | None = None,
|
|
379
|
+
encoder_fingerprint: str | None | object = _COMPUTE,
|
|
380
|
+
) -> Path:
|
|
381
|
+
"""Backbone + head + everything needed to feed them correctly.
|
|
382
|
+
|
|
383
|
+
`save_backbone` deliberately stores only the backbone, because fine-tuning replaces
|
|
384
|
+
the head. Inference needs the opposite: a head, and the encoder settings that
|
|
385
|
+
produced the vectors it was trained on. The label map is frozen because a class order
|
|
386
|
+
rebuilt at deploy time predicts confidently and wrongly; `T` and `dominant`
|
|
387
|
+
fail exactly the same way and just as silently, so they travel in the file.
|
|
388
|
+
|
|
389
|
+
`labels` is a filename, not a path — the label map ships beside the weights,
|
|
390
|
+
so the checkpoint records what to look for in its own directory rather than
|
|
391
|
+
an absolute path that will not survive being copied to the board.
|
|
392
|
+
|
|
393
|
+
`encoder_fingerprint` identifies the encoder's behaviour, so a checkpoint
|
|
394
|
+
can never be served by an encoder it was not fitted on — the same class of
|
|
395
|
+
silent, confident failure a frozen label map prevents.
|
|
396
|
+
"""
|
|
397
|
+
# Computed by default so a training run cannot forget it. Passing None
|
|
398
|
+
# explicitly writes a checkpoint without one, which is what the tests for
|
|
399
|
+
# the pre-guard path need.
|
|
400
|
+
signature = None
|
|
401
|
+
if encoder_fingerprint is _COMPUTE:
|
|
402
|
+
from islkit.features import encoder_fingerprint as compute
|
|
403
|
+
from islkit.features import encoder_signature
|
|
404
|
+
|
|
405
|
+
encoder_fingerprint = compute()
|
|
406
|
+
# Travels with the hash because the hash is exact and machines are not:
|
|
407
|
+
# a checkpoint trained on a laptop must be loadable on the board, and
|
|
408
|
+
# those two disagree in the last bit of a float32. See features.py.
|
|
409
|
+
# As a tensor, not the numpy array encoder_signature returns:
|
|
410
|
+
# load_classifier reads with weights_only=True (it should — a checkpoint
|
|
411
|
+
# arrives over scp), and that refuses a pickled numpy array outright.
|
|
412
|
+
signature = torch.from_numpy(encoder_signature())
|
|
413
|
+
|
|
414
|
+
path = Path(path)
|
|
415
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
416
|
+
torch.save(
|
|
417
|
+
{
|
|
418
|
+
"config": model.backbone.config,
|
|
419
|
+
"n_classes": model.head.out_features,
|
|
420
|
+
"labels": labels,
|
|
421
|
+
"encoder": dict(encoder),
|
|
422
|
+
"encoder_fingerprint": encoder_fingerprint,
|
|
423
|
+
"encoder_signature": signature,
|
|
424
|
+
"state_dict": model.state_dict(),
|
|
425
|
+
"training": dict(training or {}),
|
|
426
|
+
},
|
|
427
|
+
path,
|
|
428
|
+
)
|
|
429
|
+
return path
|
|
430
|
+
|
|
431
|
+
|
|
432
|
+
def load_classifier(path: str | Path, device=None) -> tuple[SignClassifier, dict]:
|
|
433
|
+
"""Rebuild a `save_classifier` file. Returns the model and its metadata."""
|
|
434
|
+
blob = torch.load(Path(path), map_location=device or "cpu", weights_only=True)
|
|
435
|
+
cfg = dict(blob["config"])
|
|
436
|
+
cfg["dilations"] = tuple(cfg["dilations"])
|
|
437
|
+
model = SignClassifier(Backbone(**cfg), blob["n_classes"])
|
|
438
|
+
model.load_state_dict(blob["state_dict"])
|
|
439
|
+
model.eval()
|
|
440
|
+
if device is not None:
|
|
441
|
+
model = model.to(device)
|
|
442
|
+
meta = {k: v for k, v in blob.items() if k != "state_dict"}
|
|
443
|
+
return model, meta
|
|
444
|
+
|
|
445
|
+
|
|
446
|
+
# --------------------------------------------------------------------------
|
|
447
|
+
# Training
|
|
448
|
+
# --------------------------------------------------------------------------
|
|
449
|
+
|
|
450
|
+
|
|
451
|
+
@dataclass
|
|
452
|
+
class TrainConfig:
|
|
453
|
+
"""One frozen recipe, for the same reason XGB_PARAMS is frozen.
|
|
454
|
+
|
|
455
|
+
The baseline was run with no tuning so that the floor it set was honest. A tuned
|
|
456
|
+
TCN compared against an untuned baseline measures the search, not the model, so these
|
|
457
|
+
are ordinary defaults chosen once and left alone.
|
|
458
|
+
|
|
459
|
+
`label_smoothing` is deliberately 0.0: the decline confidence threshold
|
|
460
|
+
has to be re-measured against this model's probabilities, and smoothing would
|
|
461
|
+
flatten exactly the distribution it needs to read.
|
|
462
|
+
"""
|
|
463
|
+
|
|
464
|
+
epochs: int = 60
|
|
465
|
+
batch: int = 64
|
|
466
|
+
lr: float = 3e-3
|
|
467
|
+
weight_decay: float = 1e-4
|
|
468
|
+
label_smoothing: float = 0.0
|
|
469
|
+
min_lr: float = 1e-5
|
|
470
|
+
warmup_epochs: int = 3
|
|
471
|
+
|
|
472
|
+
|
|
473
|
+
def _cosine_lr(epoch: int, cfg: TrainConfig) -> float:
|
|
474
|
+
"""Linear warmup then cosine decay, as a multiplier on `cfg.lr`."""
|
|
475
|
+
if epoch < cfg.warmup_epochs:
|
|
476
|
+
return (epoch + 1) / max(1, cfg.warmup_epochs)
|
|
477
|
+
span = max(1, cfg.epochs - cfg.warmup_epochs)
|
|
478
|
+
t = (epoch - cfg.warmup_epochs) / span
|
|
479
|
+
floor = cfg.min_lr / cfg.lr
|
|
480
|
+
return floor + (1 - floor) * 0.5 * (1 + math.cos(math.pi * t))
|
|
481
|
+
|
|
482
|
+
|
|
483
|
+
def fit(
|
|
484
|
+
model: SignClassifier,
|
|
485
|
+
X: np.ndarray,
|
|
486
|
+
y: np.ndarray,
|
|
487
|
+
cfg: TrainConfig | None = None,
|
|
488
|
+
device=None,
|
|
489
|
+
X_val: np.ndarray | None = None,
|
|
490
|
+
y_val: np.ndarray | None = None,
|
|
491
|
+
verbose: bool = False,
|
|
492
|
+
) -> dict[str, list[float]]:
|
|
493
|
+
"""Train in place. Returns per-epoch loss and accuracy.
|
|
494
|
+
|
|
495
|
+
Validation is optional and is only ever used for *reporting a curve*. There
|
|
496
|
+
is no early stopping and no checkpoint selection on it, because the caller
|
|
497
|
+
passes the held-out fold and selecting on that would quietly turn the
|
|
498
|
+
out-of-fold number into a fitted one.
|
|
499
|
+
"""
|
|
500
|
+
cfg = cfg or TrainConfig()
|
|
501
|
+
device = device or next(model.parameters()).device
|
|
502
|
+
|
|
503
|
+
Xt = torch.as_tensor(np.asarray(X, np.float32)).to(device)
|
|
504
|
+
yt = torch.as_tensor(np.asarray(y, np.int64)).to(device)
|
|
505
|
+
opt = torch.optim.AdamW(
|
|
506
|
+
(p for p in model.parameters() if p.requires_grad),
|
|
507
|
+
lr=cfg.lr,
|
|
508
|
+
weight_decay=cfg.weight_decay,
|
|
509
|
+
)
|
|
510
|
+
lossfn = nn.CrossEntropyLoss(label_smoothing=cfg.label_smoothing)
|
|
511
|
+
history: dict[str, list[float]] = {"loss": [], "train_acc": [], "val_acc": []}
|
|
512
|
+
|
|
513
|
+
for epoch in range(cfg.epochs):
|
|
514
|
+
for group in opt.param_groups:
|
|
515
|
+
group["lr"] = cfg.lr * _cosine_lr(epoch, cfg)
|
|
516
|
+
|
|
517
|
+
model.train()
|
|
518
|
+
order = torch.randperm(len(Xt), device=device)
|
|
519
|
+
total, correct = 0.0, 0
|
|
520
|
+
for i in range(0, len(order), cfg.batch):
|
|
521
|
+
idx = order[i : i + cfg.batch]
|
|
522
|
+
if len(idx) < 2:
|
|
523
|
+
continue # BatchNorm needs more than one row
|
|
524
|
+
opt.zero_grad()
|
|
525
|
+
logits = model(Xt[idx])
|
|
526
|
+
loss = lossfn(logits, yt[idx])
|
|
527
|
+
loss.backward()
|
|
528
|
+
opt.step()
|
|
529
|
+
total += float(loss.detach()) * len(idx)
|
|
530
|
+
correct += int((logits.argmax(1) == yt[idx]).sum())
|
|
531
|
+
|
|
532
|
+
history["loss"].append(total / len(Xt))
|
|
533
|
+
history["train_acc"].append(correct / len(Xt))
|
|
534
|
+
if X_val is not None:
|
|
535
|
+
pred = predict_proba(model, X_val, device=device).argmax(axis=1)
|
|
536
|
+
history["val_acc"].append(float((pred == np.asarray(y_val)).mean()))
|
|
537
|
+
if verbose and (epoch % 10 == 9 or epoch == 0):
|
|
538
|
+
tail = f" val {history['val_acc'][-1] * 100:5.1f}%" if X_val is not None else ""
|
|
539
|
+
print(
|
|
540
|
+
f" epoch {epoch + 1:3d} | loss {history['loss'][-1]:.4f} "
|
|
541
|
+
f"| train {history['train_acc'][-1] * 100:5.1f}%{tail}",
|
|
542
|
+
flush=True,
|
|
543
|
+
)
|
|
544
|
+
return history
|
|
545
|
+
|
|
546
|
+
|
|
547
|
+
@torch.no_grad()
|
|
548
|
+
def predict_proba(
|
|
549
|
+
model: SignClassifier, X: np.ndarray, device=None, batch: int = 256
|
|
550
|
+
) -> np.ndarray:
|
|
551
|
+
"""Class probabilities, (N, n_classes) float32 on the CPU.
|
|
552
|
+
|
|
553
|
+
Softmax rather than logits because the decline threshold and the
|
|
554
|
+
top-3 figure are both read off probabilities, and because it matches
|
|
555
|
+
`XGBBaseline.predict_proba` so both models can be driven by the same
|
|
556
|
+
cross-validation loop.
|
|
557
|
+
"""
|
|
558
|
+
device = device or next(model.parameters()).device
|
|
559
|
+
was_training = model.training
|
|
560
|
+
model.eval()
|
|
561
|
+
out = []
|
|
562
|
+
Xa = np.asarray(X, np.float32)
|
|
563
|
+
for i in range(0, len(Xa), batch):
|
|
564
|
+
chunk = torch.as_tensor(Xa[i : i + batch]).to(device)
|
|
565
|
+
out.append(torch.softmax(model(chunk), dim=1).cpu().numpy())
|
|
566
|
+
model.train(was_training)
|
|
567
|
+
return np.concatenate(out).astype(np.float32)
|
|
568
|
+
|
|
569
|
+
|
|
570
|
+
@dataclass
|
|
571
|
+
class TCNEstimator:
|
|
572
|
+
"""fit / predict_proba around the TCN, so `cross_val_predict` can drive it.
|
|
573
|
+
|
|
574
|
+
The TCN's whole claim is a same-split comparison against the baseline. Reusing its
|
|
575
|
+
cross-validation loop rather than writing a second one is what makes that
|
|
576
|
+
literally true: identical folds, identical out-of-fold accounting, and the
|
|
577
|
+
tested guarantee that no model ever scores a row it trained on.
|
|
578
|
+
"""
|
|
579
|
+
|
|
580
|
+
n_classes: int
|
|
581
|
+
cfg: TrainConfig = field(default_factory=TrainConfig)
|
|
582
|
+
seed: int = 0
|
|
583
|
+
device: object | None = None
|
|
584
|
+
backbone_kwargs: dict = field(default_factory=dict)
|
|
585
|
+
verbose: bool = False
|
|
586
|
+
# The fold this estimator will be scored on, for a validation *curve* only.
|
|
587
|
+
# `fit` does no early stopping and no checkpoint selection, so this cannot
|
|
588
|
+
# leak into the out-of-fold number — it only makes the generalisation gap
|
|
589
|
+
# visible per epoch instead of only at the end.
|
|
590
|
+
val: tuple[np.ndarray, np.ndarray] | None = None
|
|
591
|
+
model: SignClassifier | None = None
|
|
592
|
+
history: dict = field(default_factory=dict)
|
|
593
|
+
seconds: float = 0.0
|
|
594
|
+
|
|
595
|
+
def fit(self, X: np.ndarray, y: np.ndarray) -> TCNEstimator:
|
|
596
|
+
from islkit.seeding import set_seed # noqa: PLC0415 — keeps torch off the import path
|
|
597
|
+
|
|
598
|
+
set_seed(self.seed)
|
|
599
|
+
t0 = time.time()
|
|
600
|
+
self.model = build_model(self.n_classes, **self.backbone_kwargs).to(self.device)
|
|
601
|
+
X_val, y_val = self.val if self.val is not None else (None, None)
|
|
602
|
+
self.history = fit(
|
|
603
|
+
self.model,
|
|
604
|
+
X,
|
|
605
|
+
y,
|
|
606
|
+
self.cfg,
|
|
607
|
+
self.device,
|
|
608
|
+
X_val=X_val,
|
|
609
|
+
y_val=y_val,
|
|
610
|
+
verbose=self.verbose,
|
|
611
|
+
)
|
|
612
|
+
self.seconds = time.time() - t0
|
|
613
|
+
return self
|
|
614
|
+
|
|
615
|
+
@property
|
|
616
|
+
def val_size(self) -> int:
|
|
617
|
+
"""Rows in the validation curve. Checked against the fold it should be."""
|
|
618
|
+
return 0 if self.val is None else len(self.val[1])
|
|
619
|
+
|
|
620
|
+
def predict_proba(self, X: np.ndarray) -> np.ndarray:
|
|
621
|
+
if self.model is None:
|
|
622
|
+
raise RuntimeError("fit() first")
|
|
623
|
+
return predict_proba(self.model, X, device=self.device)
|