Download model.py from riteshhf/tsfp-repro-code: direct link, hf CLI and curl.
- Browser
- Download file 14.4 kB
-
https://huggingface.co/riteshhf/tsfp-repro-code/resolve/main/model.py
- Command line
-
hf download hf://riteshhf/tsfp-repro-code/model.py
-
curl -L -o model.py https://huggingface.co/riteshhf/tsfp-repro-code/resolve/main/model.py
14.4 kB
| """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) | |