TSRDA / main_method /code /temporal_srda.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
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