""" Temporal SRDA (T-SRDA) — PDAViT's Spatial-Reduction Dual-Attention (Zhou et al., Neurocomputing 2026) adapted to the TEMPORAL axis. Drop-in replacement for TemporalDaViTEncoder inside UTAE_ASPP. Matches its calling convention exactly: enc = TemporalSRDAEncoder(d_model=256, n_heads=8, R_schedule=(4, 4, 4)) out = enc(x, pos) # x: (N, T, d) pos: (N, T, d) date PE or None Mechanism per block (1D mapping of CUSA + dual attention): KV path : Linear d -> d/R ("point-wise conv") unfold: group R consecutive timesteps -> (N, T/R, d) First SA over T/R tokens (key-value "self-checking") Q path : full-resolution T tokens Second SA: cross-attention Q(T) x KV(T/R) -> GLOBAL temporal receptive field in a single layer (Swin/DaViT window schedule needs the whole [4,8,16]+shift stack to approximate this). scMLP : FC -> GELU (+ parallel depth-wise Conv1d over T) -> FC Design notes (do not change casually): * NO internal gradient checkpointing — UTAE_ASPP already checkpoints at the chunk level (_run_temporal_chunk). Layer-level checkpointing here would recreate the double-checkpoint waste found in the DaViT encoder. * Padding convention matches the existing pipeline: collate_fn zero-pads to batch max_T and the windowed encoders attend over padded steps too. We keep that convention for a clean apples-to-apples ablation. The encoder internally pads T to a multiple of max(R_schedule) and crops back, so any T works (33..61 all fine). * Date PE MUST be added before the KV reduction (handled here: x = x+pos at entry) — grouped KV tokens over irregular acquisitions are meaningless without date information. """ import torch import torch.nn as nn import torch.nn.functional as F class TemporalSRDABlock(nn.Module): def __init__(self, d_model=256, n_heads=8, R=4, mlp_ratio=4.0, attn_drop=0.0, proj_drop=0.0): super().__init__() assert d_model % R == 0, "d_model must be divisible by R" assert d_model % n_heads == 0 self.d = d_model self.R = R self.h = n_heads self.dh = d_model // n_heads self.scale = self.dh ** -0.5 # ---- KV reduction path (CUSA, 1D) --------------------------------- self.norm_kv = nn.LayerNorm(d_model) self.pw = nn.Linear(d_model, d_model // R, bias=False) # point-wise conv self.norm_first = nn.LayerNorm(d_model) self.first_sa = nn.MultiheadAttention( d_model, n_heads, dropout=attn_drop, batch_first=True) # ---- Second (dual) attention --------------------------------------- self.norm_q = nn.LayerNorm(d_model) self.q_proj = nn.Linear(d_model, d_model, bias=False) self.kv_proj = nn.Linear(d_model, 2 * d_model, bias=False) self.attn_drop = nn.Dropout(attn_drop) self.out_proj = nn.Linear(d_model, d_model) self.proj_drop = nn.Dropout(proj_drop) # ---- scMLP ---------------------------------------------------------- hidden = int(d_model * mlp_ratio) self.norm_mlp = nn.LayerNorm(d_model) self.fc1 = nn.Linear(d_model, hidden) self.dw = nn.Conv1d(hidden, hidden, kernel_size=3, padding=1, groups=hidden) # DW over T self.fc2 = nn.Linear(hidden, d_model) def forward(self, x): """x: (N, T, d) with T % R == 0 (guaranteed by the encoder wrapper).""" N, T, d = x.shape R = self.R # ---- KV path: point-wise conv -> temporal unfold -> first SA ------- z = self.pw(self.norm_kv(x)) # (N, T, d/R) z = z.reshape(N, T // R, R * (d // R)) # unfold -> (N, T/R, d) zn = self.norm_first(z) z_sa, _ = self.first_sa(zn, zn, zn, need_weights=False) kv = z + z_sa # self-checked KV # ---- second SA: full-T queries x reduced KV ------------------------- q = self.q_proj(self.norm_q(x)) k, v = self.kv_proj(kv).chunk(2, dim=-1) q = q.reshape(N, T, self.h, self.dh).transpose(1, 2) # (N,h,T,dh) k = k.reshape(N, -1, self.h, self.dh).transpose(1, 2) # (N,h,T/R,dh) v = v.reshape(N, -1, self.h, self.dh).transpose(1, 2) attn = (q @ k.transpose(-2, -1)) * self.scale # (N,h,T,T/R) attn = self.attn_drop(attn.softmax(dim=-1)) out = (attn @ v).transpose(1, 2).reshape(N, T, d) x = x + self.proj_drop(self.out_proj(out)) # ---- scMLP ----------------------------------------------------------- y = self.fc1(self.norm_mlp(x)) y = F.gelu(y) + self.dw(y.transpose(1, 2)).transpose(1, 2) x = x + self.fc2(y) return x class TemporalSRDAEncoder(nn.Module): """ Signature-compatible replacement for TemporalDaViTEncoder. forward(x, pos=None): x (N, T, d), pos (N, T, d) date positional encoding. """ def __init__(self, d_model=256, n_heads=8, R_schedule=(4, 4, 4), **kwargs): super().__init__() # **kwargs silently absorbs legacy windows=/shifts= if passed through self.R_max = max(R_schedule) self.blocks = nn.ModuleList( TemporalSRDABlock(d_model, n_heads, R=r) for r in R_schedule) def forward(self, x, pos=None): if pos is not None: x = x + pos # date PE BEFORE KV reduction N, T, d = x.shape pad = (-T) % self.R_max # pad T up to multiple of R_max if pad: x = F.pad(x, (0, 0, 0, pad)) # zero-pad — matches pipeline for blk in self.blocks: x = blk(x) if pad: x = x[:, :T] return x