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/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)