Spaces:
Running on Zero
Running on Zero
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,
)
|