tsfp-repro-code / model.py
riteshhf's picture
Upload folder using huggingface_hub
2188a91 verified
Raw
History Blame Contribute Delete
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)