File size: 5,831 Bytes
71d64bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
"""
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