| """Independent re-implementation of TS-Fingerprint (ICML 2026, OpenReview lrRBHFIgaK). |
| |
| No official code was released for the paper, so every component here is derived |
| from the paper text: Sec 3.2 (architecture), Sec 3.3 (losses), Sec 3.4 (attention |
| pooling head) and Sec 4.1 (hyperparameters: 6-layer encoder, 2-layer decoder, |
| d=128, 8 heads, k=8, mask ratio 0.6, lambda=1e-4). |
| """ |
|
|
| import math |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
|
|
| class PatchEmbed(nn.Module): |
| """Non-overlapping patching over time; all C channels enter the patch (Medformer-style).""" |
|
|
| def __init__(self, c_in, patch_size, d_model): |
| super().__init__() |
| self.patch_size = patch_size |
| self.proj = nn.Linear(c_in * patch_size, d_model) |
|
|
| def forward(self, x): |
| b, t, c = x.shape |
| n = t // self.patch_size |
| x = x[:, : n * self.patch_size].reshape(b, n, self.patch_size * c) |
| return self.proj(x) |
|
|
|
|
| class CrossAttnBlock(nn.Module): |
| """One iteration of the Perceiver-style bottleneck: latents cross-attend to the |
| patch sequence, then self-attend among themselves (Sec 3.2, 'iterative cross-attention').""" |
|
|
| def __init__(self, d_model, n_heads, dropout=0.1): |
| super().__init__() |
| self.ln_q1 = nn.LayerNorm(d_model) |
| self.ln_kv = nn.LayerNorm(d_model) |
| self.cross = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True) |
| self.ln_q2 = nn.LayerNorm(d_model) |
| self.self_attn = nn.MultiheadAttention(d_model, n_heads, dropout=dropout, batch_first=True) |
| self.ln_q3 = nn.LayerNorm(d_model) |
| self.ff = nn.Sequential( |
| nn.Linear(d_model, 4 * d_model), nn.GELU(), nn.Dropout(dropout), |
| nn.Linear(4 * d_model, d_model), nn.Dropout(dropout), |
| ) |
|
|
| def forward(self, q, kv, kv_mask=None, need_weights=False): |
| h, attn = self.cross(self.ln_q1(q), self.ln_kv(kv), self.ln_kv(kv), |
| key_padding_mask=kv_mask, need_weights=need_weights, |
| average_attn_weights=True) |
| q = q + h |
| h, _ = self.self_attn(self.ln_q2(q), self.ln_q2(q), self.ln_q2(q), need_weights=False) |
| q = q + h |
| q = q + self.ff(self.ln_q3(q)) |
| return q, attn |
|
|
|
|
| class FingerprintEncoder(nn.Module): |
| """E_theta : X -> F' in R^{k x d}. Fixed-rank bottleneck via a learnable query set |
| Q in R^{k x d} with k << T (Claim 1).""" |
|
|
| def __init__(self, c_in, patch_size, d_model=128, n_heads=8, n_layers=6, k=8, |
| max_patches=512, dropout=0.1): |
| super().__init__() |
| self.k = k |
| self.d_model = d_model |
| self.patch_embed = PatchEmbed(c_in, patch_size, d_model) |
| self.pos = nn.Parameter(torch.zeros(1, max_patches, d_model)) |
| nn.init.trunc_normal_(self.pos, std=0.02) |
| |
| self.Q = nn.Parameter(torch.randn(k, d_model) * 0.02) |
| self.blocks = nn.ModuleList( |
| [CrossAttnBlock(d_model, n_heads, dropout) for _ in range(n_layers)] |
| ) |
| self.norm = nn.LayerNorm(d_model) |
|
|
| def embed_patches(self, x): |
| p = self.patch_embed(x) |
| return p + self.pos[:, : p.shape[1]] |
|
|
| def forward(self, x, kv_mask=None, return_attn=False): |
| """x: (B, T, C) -> F': (B, k, d). Output shape is independent of T.""" |
| p = self.embed_patches(x) |
| q = self.Q.unsqueeze(0).expand(x.shape[0], -1, -1) |
| attns = [] |
| for blk in self.blocks: |
| q, a = blk(q, p, kv_mask=kv_mask, need_weights=return_attn) |
| if return_attn: |
| attns.append(a) |
| f = self.norm(q) |
| if return_attn: |
| return f, attns |
| return f |
|
|
|
|
| class FingerprintDecoder(nn.Module): |
| """D_phi conditions *solely* on F'. Mask tokens carry only positional information, |
| so the data-processing chain X -> F' -> X_hat is strict (Sec 3.2, Decoder).""" |
|
|
| def __init__(self, c_in, patch_size, d_model=128, n_heads=8, n_layers=2, |
| max_patches=512, dropout=0.1): |
| super().__init__() |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, d_model)) |
| self.pos = nn.Parameter(torch.zeros(1, max_patches, d_model)) |
| nn.init.trunc_normal_(self.pos, std=0.02) |
| self.blocks = nn.ModuleList( |
| [CrossAttnBlock(d_model, n_heads, dropout) for _ in range(n_layers)] |
| ) |
| self.norm = nn.LayerNorm(d_model) |
| self.head = nn.Linear(d_model, c_in * patch_size) |
|
|
| def forward(self, f, n_patches): |
| b = f.shape[0] |
| t = self.mask_token.expand(b, n_patches, -1) + self.pos[:, :n_patches] |
| for blk in self.blocks: |
| t, _ = blk(t, f) |
| return self.head(self.norm(t)) |
|
|
|
|
| def total_coding_rate_loss(f, eps=0.5): |
| """L_div = -1/2 log det(I + d/(k eps^2) F'^T F') (Eq. 3). |
| |
| Computed per sample over its k tokens; by Sylvester's identity the d x d and |
| k x k forms have identical value, and we use the cheaper k x k Gram form. |
| """ |
| b, k, d = f.shape |
| f = f.float() |
| fn = F.normalize(f, dim=-1) |
| gram = torch.bmm(fn, fn.transpose(1, 2)) |
| ident = torch.eye(k, device=f.device, dtype=f.dtype).unsqueeze(0) |
| mat = ident + (d / (k * eps ** 2)) * gram |
| return -0.5 * torch.linalg.slogdet(mat)[1].mean() |
|
|
|
|
| class TSFingerprint(nn.Module): |
| def __init__(self, c_in, patch_size, n_classes, d_model=128, n_heads=8, |
| enc_layers=6, dec_layers=2, k=8, max_patches=512, dropout=0.1): |
| super().__init__() |
| self.patch_size = patch_size |
| self.encoder = FingerprintEncoder(c_in, patch_size, d_model, n_heads, |
| enc_layers, k, max_patches, dropout) |
| self.decoder = FingerprintDecoder(c_in, patch_size, d_model, n_heads, |
| dec_layers, max_patches, dropout) |
| |
| self.q_task = nn.Linear(d_model, 1, bias=False) |
| self.cls_head = nn.Sequential(nn.LayerNorm(d_model), nn.Linear(d_model, n_classes)) |
|
|
| |
| def patchify(self, x): |
| b, t, c = x.shape |
| n = t // self.patch_size |
| return x[:, : n * self.patch_size].reshape(b, n, self.patch_size * c) |
|
|
| def pretrain_step(self, x, mask_ratio=0.6, lam=1e-4, use_div=True, generator=None): |
| target = self.patchify(x) |
| b, n, _ = target.shape |
| |
| n_keep = max(1, int(round(n * (1 - mask_ratio)))) |
| noise = torch.rand(b, n, device=x.device, generator=generator) |
| keep = noise.argsort(dim=1)[:, :n_keep] |
| kv_mask = torch.ones(b, n, dtype=torch.bool, device=x.device) |
| kv_mask.scatter_(1, keep, False) |
| f = self.encoder(x, kv_mask=kv_mask) |
| pred = self.decoder(f, n) |
| l_rec = F.mse_loss(pred, target) |
| l_div = total_coding_rate_loss(f) if use_div else torch.zeros((), device=x.device) |
| return l_rec + lam * l_div, l_rec.detach(), l_div.detach() |
|
|
| |
| def forward(self, x, return_alpha=False): |
| f = self.encoder(x) |
| alpha = torch.softmax(self.q_task(f).squeeze(-1), dim=-1) |
| z = (alpha.unsqueeze(-1) * f).sum(1) |
| logits = self.cls_head(z) |
| if return_alpha: |
| return logits, alpha, f |
| return logits |
|
|
|
|
| |
| |
| |
| class PlainEncoder(nn.Module): |
| """Standard transformer patch encoder -> variable-length token sequence |
| (the 'entangled view' shared by Ti-MAE and SimMTM).""" |
|
|
| def __init__(self, c_in, patch_size, d_model=128, n_heads=8, n_layers=6, |
| max_patches=512, dropout=0.1): |
| super().__init__() |
| self.patch_embed = PatchEmbed(c_in, patch_size, d_model) |
| self.pos = nn.Parameter(torch.zeros(1, max_patches, d_model)) |
| nn.init.trunc_normal_(self.pos, std=0.02) |
| layer = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, |
| activation="gelu", batch_first=True, |
| norm_first=True) |
| self.enc = nn.TransformerEncoder(layer, n_layers) |
| self.norm = nn.LayerNorm(d_model) |
|
|
| def forward(self, x, src_key_padding_mask=None): |
| p = self.patch_embed(x) |
| p = p + self.pos[:, : p.shape[1]] |
| return self.norm(self.enc(p, src_key_padding_mask=src_key_padding_mask)) |
|
|
|
|
| class TiMAE(nn.Module): |
| """Ti-MAE (Li et al., 2023): masked patch autoencoding on a plain transformer, |
| visible patches are fed to the encoder, mask tokens are appended in the decoder, |
| downstream head is global average pooling.""" |
|
|
| def __init__(self, c_in, patch_size, n_classes, d_model=128, n_heads=8, |
| enc_layers=6, dec_layers=2, max_patches=512, dropout=0.1): |
| super().__init__() |
| self.patch_size = patch_size |
| self.encoder = PlainEncoder(c_in, patch_size, d_model, n_heads, enc_layers, |
| max_patches, dropout) |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, d_model)) |
| self.dec_pos = nn.Parameter(torch.zeros(1, max_patches, d_model)) |
| nn.init.trunc_normal_(self.dec_pos, std=0.02) |
| layer = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, |
| activation="gelu", batch_first=True, |
| norm_first=True) |
| self.dec = nn.TransformerEncoder(layer, dec_layers) |
| self.head_rec = nn.Linear(d_model, c_in * patch_size) |
| self.cls_head = nn.Sequential(nn.LayerNorm(d_model), nn.Linear(d_model, n_classes)) |
|
|
| def patchify(self, x): |
| b, t, c = x.shape |
| n = t // self.patch_size |
| return x[:, : n * self.patch_size].reshape(b, n, self.patch_size * c) |
|
|
| def pretrain_step(self, x, mask_ratio=0.6, generator=None, **kw): |
| target = self.patchify(x) |
| b, n, _ = target.shape |
| n_keep = max(1, int(round(n * (1 - mask_ratio)))) |
| noise = torch.rand(b, n, device=x.device, generator=generator) |
| keep = noise.argsort(dim=1)[:, :n_keep] |
| pad = torch.ones(b, n, dtype=torch.bool, device=x.device) |
| pad.scatter_(1, keep, False) |
| h = self.encoder(x, src_key_padding_mask=pad) |
| |
| h = torch.where(pad.unsqueeze(-1), self.mask_token.expand(b, n, -1), h) |
| h = h + self.dec_pos[:, :n] |
| pred = self.head_rec(self.dec(h)) |
| l_rec = F.mse_loss(pred, target) |
| return l_rec, l_rec.detach(), torch.zeros((), device=x.device) |
|
|
| def forward(self, x): |
| h = self.encoder(x) |
| return self.cls_head(h.mean(1)) |
|
|
|
|
| class SimMTM(nn.Module): |
| """SimMTM (Dong et al., 2023): reconstruct the original series from *multiple* |
| masked views, aggregating them by point-wise series-similarity, plus a |
| series-wise contrastive term. Head is global average pooling.""" |
|
|
| def __init__(self, c_in, patch_size, n_classes, d_model=128, n_heads=8, |
| enc_layers=6, dec_layers=2, max_patches=512, dropout=0.1, |
| n_views=3, temperature=0.2): |
| super().__init__() |
| self.patch_size = patch_size |
| self.n_views = n_views |
| self.temperature = temperature |
| self.encoder = PlainEncoder(c_in, patch_size, d_model, n_heads, enc_layers, |
| max_patches, dropout) |
| layer = nn.TransformerEncoderLayer(d_model, n_heads, 4 * d_model, dropout, |
| activation="gelu", batch_first=True, |
| norm_first=True) |
| self.dec = nn.TransformerEncoder(layer, dec_layers) |
| self.head_rec = nn.Linear(d_model, c_in * patch_size) |
| self.proj = nn.Sequential(nn.Linear(d_model, d_model), nn.GELU(), |
| nn.Linear(d_model, d_model)) |
| self.cls_head = nn.Sequential(nn.LayerNorm(d_model), nn.Linear(d_model, n_classes)) |
|
|
| def patchify(self, x): |
| b, t, c = x.shape |
| n = t // self.patch_size |
| return x[:, : n * self.patch_size].reshape(b, n, self.patch_size * c) |
|
|
| def pretrain_step(self, x, mask_ratio=0.6, generator=None, **kw): |
| target = self.patchify(x) |
| b, n, _ = target.shape |
| views = [] |
| for _ in range(self.n_views): |
| m = (torch.rand(b, x.shape[1], 1, device=x.device, generator=generator) |
| < mask_ratio) |
| views.append(x.masked_fill(m, 0.0)) |
| xs = torch.cat(views, 0) |
| h = self.encoder(xs) |
| s = self.proj(h.mean(1)) |
| s = F.normalize(s, dim=-1) |
| |
| sim = (s.view(self.n_views, b, -1) * s.view(self.n_views, b, -1)[0:1]).sum(-1) |
| w = torch.softmax(sim / self.temperature, dim=0) |
| agg = (h.view(self.n_views, b, n, -1) * w[..., None, None]).sum(0) |
| pred = self.head_rec(self.dec(agg)) |
| l_rec = F.mse_loss(pred, target) |
| |
| logits = s @ s.t() / self.temperature |
| lbl = torch.arange(b, device=x.device).repeat(self.n_views) |
| eye = torch.eye(logits.shape[0], device=x.device, dtype=torch.bool) |
| logits = logits.masked_fill(eye, -1e4) |
| pos = (lbl[:, None] == lbl[None, :]) & ~eye |
| log_p = torch.log_softmax(logits, dim=-1) |
| l_con = -(log_p * pos).sum(-1).div(pos.sum(-1).clamp(min=1)).mean() |
| return l_rec + 0.1 * l_con, l_rec.detach(), l_con.detach() |
|
|
| def forward(self, x): |
| h = self.encoder(x) |
| return self.cls_head(h.mean(1)) |
|
|
|
|
| MODELS = {"tsfp": TSFingerprint, "timae": TiMAE, "simmtm": SimMTM} |
|
|
|
|
| def build(name, **kw): |
| if name == "simmtm": |
| kw.pop("k", None) |
| elif name == "timae": |
| kw.pop("k", None) |
| return MODELS[name](**kw) |
|
|