Download main_method/code/temporal_srda.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 5.83 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/temporal_srda.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/temporal_srda.py
-
curl -L -o temporal_srda.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/temporal_srda.py
5.83 kB
| """ | |
| 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 | |