File size: 9,807 Bytes
87608ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
"""Context-Modulated Pixel Transformer (CM-PiT) building blocks.

CM-PiT compresses local dense pixel features into attention tokens, processes
them with gated self-attention and SwiGLU, and expands them back without losing
the original pixel lattice. Global encoder tokens generate adaptive shift,
scale, and residual gates that condition both transformer sublayers.
"""

from typing import Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange

from .Gated_Attention import GatedAttention
from .RoPE import RotaryPositionEmbedding2D
from .precision import full_precision


def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
    """Apply affine context modulation without changing the tensor layout.

    Args:
        x: Normalized pixel tokens with shape ``[B, N, P, C]``.
        shift: Context-predicted additive offsets with shape ``[B, N, P, C]``.
        scale: Context-predicted residual scales with shape ``[B, N, P, C]``.

    Returns:
        Modulated tokens ``x * (1 + scale) + shift`` with shape
        ``[B, N, P, C]``.
    """
    return x * (1.0 + scale) + shift


class SwiGLU(nn.Module):
    """SwiGLU feed-forward layer operating independently on every pixel token.

    The first projection creates value and gate branches, SiLU activates the
    gate, and the second projection returns to the pixel-channel dimension.
    Spatial and patch axes are preserved throughout the module.
    """

    def __init__(self, dim: int, hidden_dim: int) -> None:
        """Construct the gated feed-forward projections.

        Args:
            dim: Input and output channel count ``C``.
            hidden_dim: Width of each hidden value/gate branch.

        Returns:
            ``None``. Learnable linear layers are registered on the module.
        """
        super().__init__()
        self.fc1 = nn.Linear(dim, hidden_dim * 2)
        self.fc2 = nn.Linear(hidden_dim, dim)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """Transform pixel tokens with a SiLU-gated hidden representation.

        Args:
            x: Floating tensor with arbitrary leading dimensions and final
                channel dimension ``C=dim``. CM-PiT supplies ``[B,N,P,C]``.

        Returns:
            Tensor with the same shape and dtype as ``x``.
        """
        value, gate = self.fc1(x).chunk(2, dim=-1)
        return self.fc2(value * F.silu(gate))


class ContextAdaNorm(nn.Module):
    """Predict Context-Guided Adaptive Normalization parameters.

    Each global context token produces shift, scale, and residual-gate values
    for both the attention and MLP sublayers over every pixel represented by
    that encoder token. The six parameter groups are unpacked by
    :class:`CMPiTBlock`.
    """

    def __init__(self, dim_ctx: int, patch_size: int, dim_pix: int) -> None:
        """Create the context-to-modulation projection.

        Args:
            dim_ctx: Channel count of each Global Context Encoder token.
            patch_size: Encoder patch side length ``P_ctx`` in image pixels.
            dim_pix: Pixel-feature channel count ``C_pix``.

        Returns:
            ``None``. The projection outputs ``6 * P_ctx^2 * C_pix`` values per
            context token.
        """
        super().__init__()
        self.proj = nn.Sequential(
            nn.SiLU(),
            nn.Linear(dim_ctx, 6 * patch_size * patch_size * dim_pix),
        )

    def forward(self, ctx: torch.Tensor) -> torch.Tensor:
        """Project context tokens into six dense pixel-wise parameter fields.

        Args:
            ctx: Context token tensor ``[B, N_ctx, C_ctx]``.

        Returns:
            Modulation tensor ``[B, N_ctx, 6 * P_ctx^2 * C_pix]``.
        """
        return self.proj(ctx)


class CMPiTBlock(nn.Module):
    """Context-Modulated Pixel Transformer block.

    The block groups a dense pixel feature map into local patches, linearly
    compresses every patch to an attention token, applies gated global
    self-attention, expands the token back to pixel features, and follows it
    with a per-pixel SwiGLU MLP. Both residual branches use Context-Guided
    Adaptive Normalization generated from DINO context tokens.
    """

    def __init__(
        self,
        dim_ctx: int,
        ctx_patch_size: int,
        dim_pix: int,
        patch_size: int,
        attn_dim: int,
        num_heads: int,
        mlp_ratio: float = 4.0,
        qk_norm: bool = True,
        rope: Optional[RotaryPositionEmbedding2D] = None,
        eps: float = 1e-6,
    ) -> None:
        """Configure one CM-PiT block.

        Args:
            dim_ctx: Context-token channel count ``C_ctx``.
            ctx_patch_size: Image patch size ``P_ctx`` represented by one
                context token.
            dim_pix: Dense pixel-feature channel count ``C_pix``.
            patch_size: Side length ``P`` grouped into one attention token.
                It must divide ``ctx_patch_size``.
            attn_dim: Compressed attention-token channel count ``D``.
            num_heads: Number of attention heads. ``D`` must be divisible by it.
            mlp_ratio: Expansion ratio controlling the SwiGLU hidden width.
            qk_norm: Whether to apply FP32 RMSNorm to each query/key head.
            rope: Optional 2D rotary position embedding shared by decoder blocks.
            eps: Numerical epsilon used by RMSNorm layers.

        Returns:
            ``None``. Attention, modulation, MLP, and projection layers are
            registered on the block.
        """
        super().__init__()
        if ctx_patch_size % patch_size != 0:
            raise ValueError(
                f"ctx_patch_size ({ctx_patch_size}) must be divisible by patch_size ({patch_size})"
            )

        self.dim_ctx = dim_ctx
        self.dim_pix = dim_pix
        self.ctx_patch_size = ctx_patch_size
        self.patch_size = patch_size
        patch_dim = patch_size * patch_size * dim_pix

        self.norm1 = nn.RMSNorm(dim_pix, eps=eps)
        self.linear_compress = nn.Linear(patch_dim, attn_dim)
        self.attn = GatedAttention(attn_dim, num_heads, qk_norm=qk_norm, rope=rope, eps=eps)
        self.linear_expand = nn.Linear(attn_dim, patch_dim)
        self.norm2 = nn.RMSNorm(dim_pix, eps=eps)
        hidden_dim = max(1, int(round(dim_pix * mlp_ratio * 2.0 / 3.0)))
        self.mlp = SwiGLU(dim_pix, hidden_dim)
        self.ada_norm = ContextAdaNorm(dim_ctx, ctx_patch_size, dim_pix)

    @staticmethod
    def _norm(norm: nn.Module, x: torch.Tensor) -> torch.Tensor:
        """Evaluate a normalization layer in FP32 and restore input dtype.

        Args:
            norm: Normalization module acting on the final channel dimension.
            x: Pixel tokens ``[B, N, P^2, C_pix]`` in the active model dtype.

        Returns:
            Normalized tensor with the same shape and dtype as ``x``.
        """
        dtype = x.dtype
        with full_precision(x.device):
            out = norm(x.float())
        return out.to(dtype)

    def _modulation(self, ctx: torch.Tensor, height: int, width: int) -> torch.Tensor:
        """Align context modulation fields with the block's pixel patches.

        Args:
            ctx: Global context tokens ``[B, H_ctx*W_ctx, C_ctx]``.
            height: Dense pixel-map height ``H``.
            width: Dense pixel-map width ``W``.

        Returns:
            Six modulation groups with shape
            ``[B, (H/P)*(W/P), 6, P^2, C_pix]``. Rearrangement is exact and
            contains no interpolation.
        """
        batch = ctx.shape[0]
        p_ctx, p = self.ctx_patch_size, self.patch_size
        ctx_h, ctx_w = height // p_ctx, width // p_ctx
        if ctx.shape[1] != ctx_h * ctx_w:
            raise ValueError(
                f"Context token count ({ctx.shape[1]}) does not match grid ({ctx_h}x{ctx_w})"
            )

        mod = self.ada_norm(ctx).view(batch, ctx_h, ctx_w, 6, p_ctx, p_ctx, self.dim_pix)
        if p_ctx == p:
            return rearrange(mod, "b h w m ph pw c -> b (h w) m (ph pw) c")

        ratio = p_ctx // p
        return rearrange(
            mod,
            "b h w m (rh ph) (rw pw) c -> b (h rh w rw) m (ph pw) c",
            rh=ratio,
            rw=ratio,
            ph=p,
            pw=p,
        )

    def forward(self, x: torch.Tensor, ctx: torch.Tensor, pos: torch.Tensor) -> torch.Tensor:
        """Apply context-modulated attention and MLP residual updates.

        Args:
            x: Dense pixel features ``[B, C_pix, H, W]``.
            ctx: Global context tokens ``[B, (H/P_ctx)*(W/P_ctx), C_ctx]``.
            pos: Integer 2D token positions ``[B, (H/P)*(W/P), 2]`` used by
                rotary position embedding in self-attention.

        Returns:
            Updated dense pixel features ``[B, C_pix, H, W]``.
        """
        batch, _, height, width = x.shape
        p = self.patch_size
        pix = rearrange(x, "b c (h ph) (w pw) -> b (h w) (ph pw) c", ph=p, pw=p)
        shift_attn, scale_attn, gate_attn, shift_mlp, scale_mlp, gate_mlp = self._modulation(
            ctx, height, width
        ).unbind(dim=2)

        out = modulate(self._norm(self.norm1, pix), shift_attn, scale_attn)
        out = self.linear_compress(out.flatten(2))
        out = self.attn(out, pos=pos)
        out = self.linear_expand(out).view(batch, -1, p * p, self.dim_pix)
        pix = pix + gate_attn * out

        out = modulate(self._norm(self.norm2, pix), shift_mlp, scale_mlp)
        pix = pix + gate_mlp * self.mlp(out)
        return rearrange(
            pix,
            "b (h w) (ph pw) c -> b c (h ph) (w pw)",
            h=height // p,
            w=width // p,
            ph=p,
            pw=p,
        )