"""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): # x: (B, T, C) 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) # (B, N, d) 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) # the learnable query set Q 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)) # (B, N, patch*C) 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() # slogdet is not fp16-safe; keep the geometric term in fp32 fn = F.normalize(f, dim=-1) # fixed-energy constraint of Lemma 3.2 gram = torch.bmm(fn, fn.transpose(1, 2)) # (B, k, k) 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) # Sec 3.4: attention pooling over the k fingerprint tokens 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)) # ---- pre-training ---- 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 # masked view: keep (1-r) of the patches visible to the encoder 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) # True == ignore 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() # ---- downstream ---- def forward(self, x, return_alpha=False): f = self.encoder(x) alpha = torch.softmax(self.q_task(f).squeeze(-1), dim=-1) # (B, k) z = (alpha.unsqueeze(-1) * f).sum(1) logits = self.cls_head(z) if return_alpha: return logits, alpha, f return logits # ------------------------------------------------------------------ # Baselines for Claim 4 # ------------------------------------------------------------------ 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) # replace masked positions with the mask token, then decode the full sequence 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) # (V*B, N, d) s = self.proj(h.mean(1)) # series-level embedding s = F.normalize(s, dim=-1) # point-wise aggregation weighted by series similarity to view 0 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) # (V, B) 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) # series-wise contrastive: views of the same series are positives 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)