File size: 16,643 Bytes
f065e53 eac5498 f065e53 eac5498 f065e53 | 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 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 | """SpatialDiffuseSlot — Semanticist arch + tok_L init + OUR spatial conditioning.
Design (2026-07-15, user-confirmed):
- encoder : Semanticist ViT (vit_base_patch16, tok_L init, num_slots=256 kept so
the ckpt loads 100%). We IGNORE its slots and take the 256 PATCH
tokens (patches never attend slots -> patch path is clean).
- pool : OUR AttnPool port (per-level learnable queries cross-attend the
16x16 patch grid) -> 85 multi-res tokens (8x8+4x4+2x2+1x1).
- dit : Semanticist DiT-L (tok_L init; null_cond re-init at K=85) with
OUR spatial-align mask applied inside its concat self-attention:
latent->latent full | latent->cond xa-rule | cond->latent BLOCK
| cond->cond identity (see spatial_mask.py, unit-tested)
- warmup : freeze_dit=True freezes the pretrained DiT trunk (blocks/embedders/
final_layer) so encoder+pool+cond-embedder adapt first; resume with
freeze_dit=False (+low lr) to open the trunk. Preserves tok_L init.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from semanticist.stage1.diffuse_slot import DiffuseSlot, DiT_with_autoenc_cond
from semanticist.stage1 import vision_transformer
import sys, os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
from spatial_mask import build_concat_self_attn_mask
def _sincos_2d(dim, s):
"""(s*s, dim) 2D sin-cos grid embedding (query localization head-start)."""
yy, xx = torch.meshgrid(torch.arange(s), torch.arange(s), indexing="ij")
grid = torch.stack([yy, xx], 0).float().reshape(2, -1) # (2, s*s)
d4 = dim // 4
omega = 1.0 / (10000 ** (torch.arange(d4).float() / d4))
out = []
for g in grid: # y then x
ang = g[:, None] * omega[None]
out += [torch.sin(ang), torch.cos(ang)]
emb = torch.cat(out, dim=1) # (s*s, 4*d4)
if emb.shape[1] < dim:
emb = torch.cat([emb, torch.zeros(emb.shape[0], dim - emb.shape[1])], 1)
return emb
class LevelNestedSampler(nn.Module):
"""Nested level-drop for our multi-res cond tokens (OUR method, port).
Cond token order is FINE-first: [8x8(0..63), 4x4(64..79), 2x2(80..83), 1x1(84)].
We drop COARSE-to-FINE keep order: keep 1x1 first, drop 8x8 first (finest, and
8x8 is 1:1 with DiT image tokens so it's the pure per-patch detail level).
Drop is WHOLE-LEVEL, coarse-first: keep_levels ~ uniform(1, num_levels) so each of
{1x1}, {1x1,2x2}, {1x1,2x2,4x4}, {all} is 25% likely (num_levels=4). This trains
every granularity equally (coarse-only levels are properly learned, unlike a
token-proportional cut which almost never fully drops 8x8). 1x1(global) always kept
-- full drop is CFG's job (separate uncond_drop). Returns (B, M) bool, True = KEEP.
inference_with_n_slots: if given, interpreted as a TOKEN count and mapped to the
largest whole keep_levels whose token total <= it (so trainer's test_num_slots=85
-> all levels; smaller -> coarse prefix).
"""
def __init__(self, level_sizes=(8, 4, 2, 1)):
super().__init__()
self.M = sum(s * s for s in level_sizes)
self.num_levels = len(level_sizes)
# coarse-first levels and cumulative token counts
coarse = sorted(level_sizes) # [1,2,4,8]
off, base = {}, 0
for s in level_sizes: # fine-first layout offsets
off[s] = base; base += s * s
# ranks: token -> which coarse-first LEVEL it belongs to (0=coarsest 1x1)
ranks = torch.empty(self.M, dtype=torch.long)
cum = []
acc = 0
for lvl, s in enumerate(coarse):
for t in range(off[s], off[s] + s * s):
ranks[t] = lvl
acc += s * s; cum.append(acc) # cumulative tokens keeping lvl+1 levels
self.register_buffer("level_of_token", ranks) # (M,) level index per token
self.register_buffer("cum_tokens", torch.tensor(cum)) # [1,5,21,85]
self.num_slots = self.M
def forward(self, batch_size, device, inference_with_n_slots=-1):
if self.training:
keep_levels = torch.randint(1, self.num_levels + 1, (batch_size,), device=device)
else:
if inference_with_n_slots != -1:
# map a token budget -> largest whole level-prefix that fits
cum = self.cum_tokens.to(device)
kl = int((cum <= inference_with_n_slots).sum().clamp(min=1))
keep_levels = torch.full((batch_size,), kl, device=device)
else:
keep_levels = torch.full((batch_size,), self.num_levels, device=device)
# keep token if its coarse-first level index < keep_levels
return self.level_of_token.to(device)[None, :] < keep_levels[:, None] # (B,M) True=keep
class SpatialAttnPool(nn.Module):
"""Port of OUR AttnPool: per-level queries cross-attend shared patch K/V.
Per level s in level_sizes: s*s learnable queries (2D sin-cos init).
depth stacked (cross-attn + FFN) layers, K/V shared across levels/layers.
Output: (B, sum(s*s), enc_d) in level order given (finest first to match
our cond-token order: 8x8, 4x4, 2x2, 1x1 -> 85 tokens).
"""
def __init__(self, enc_d=768, level_sizes=(8, 4, 2, 1), num_heads=12,
depth=2, use_ffn=True):
super().__init__()
self.level_sizes = tuple(level_sizes)
self.enc_d, self.h = enc_d, num_heads
self.hd = enc_d // num_heads
self.depth, self.use_ffn = depth, use_ffn
self.kv = nn.Linear(enc_d, 2 * enc_d)
self.kv_norm = nn.LayerNorm(enc_d)
self.queries = nn.ParameterList()
self.q_projs = nn.ModuleList()
self.out_projs = nn.ModuleList()
self.q_norms = nn.ModuleList()
self.ffns = nn.ModuleList()
for s in self.level_sizes:
q = nn.Parameter(_sincos_2d(enc_d, s) * 0.02)
self.queries.append(q)
self.q_projs.append(nn.ModuleList(
[nn.Linear(enc_d, enc_d) for _ in range(depth)]))
self.out_projs.append(nn.ModuleList(
[nn.Linear(enc_d, enc_d) for _ in range(depth)]))
self.q_norms.append(nn.ModuleList(
[nn.LayerNorm(enc_d) for _ in range(depth)]))
self.ffns.append(nn.ModuleList([
nn.Sequential(nn.LayerNorm(enc_d), nn.Linear(enc_d, 4 * enc_d),
nn.GELU(), nn.Linear(4 * enc_d, enc_d))
if use_ffn else nn.Identity() for _ in range(depth)]))
def forward(self, feats): # feats: (B, P, enc_d) patch tokens
B, P, D = feats.shape
kv = self.kv(self.kv_norm(feats)).reshape(B, P, 2, self.h, self.hd)
k = kv[:, :, 0].transpose(1, 2) # (B,h,P,hd)
v = kv[:, :, 1].transpose(1, 2)
outs = []
for li, s in enumerate(self.level_sizes):
q = self.queries[li].unsqueeze(0).expand(B, -1, -1) # (B,s*s,D)
for d in range(self.depth):
qh = self.q_projs[li][d](self.q_norms[li][d](q))
qh = qh.reshape(B, s * s, self.h, self.hd).transpose(1, 2)
att = F.scaled_dot_product_attention(qh, k, v) # (B,h,s*s,hd)
att = att.transpose(1, 2).reshape(B, s * s, D)
q = q + self.out_projs[li][d](att)
if self.use_ffn:
q = q + self.ffns[li][d](q)
outs.append(q)
return torch.cat(outs, dim=1) # (B, 85, D)
class DiTSpatial(DiT_with_autoenc_cond):
"""Semanticist DiT + our spatial mask in the concat self-attention.
forward/forward_with_cfg mirror the parent but pass `mask` to every block."""
def __init__(self, *args, level_sizes=(8, 4, 2, 1), **kwargs):
super().__init__(*args, **kwargs)
G = self.x_embedder.grid_size[0] if hasattr(self.x_embedder, "grid_size") \
else int(self.x_embedder.num_patches ** 0.5)
mask = build_concat_self_attn_mask(G, level_sizes) # (N+M, N+M) bool
self.register_buffer("spatial_mask", mask, persistent=False)
def forward(self, x, t, autoenc_cond, drop_mask=None):
x = self.x_embedder(x) + self.pos_embed
num_tokens = x.shape[1]
c = self.t_embedder(t)
autoenc = self.embed_cond(autoenc_cond, drop_mask)
x = torch.cat((x, autoenc), dim=1)
assert x.shape[1] == self.spatial_mask.shape[0], \
f"seq {x.shape[1]} != mask {self.spatial_mask.shape[0]}"
for i, block in enumerate(self.blocks):
if (i + 1) == self.encoder_depth and self.use_repa:
projected = self.projector(x)
self._repa_hook = projected[:, :num_tokens]
x = block(x, c, self.spatial_mask)
x = x[:, :num_tokens]
x = self.final_layer(x, c)
return self.unpatchify(x)
def forward_with_cfg(self, x, t, autoenc_cond, drop_mask=None, y=None, cfg_scale=1.0):
# parent's CFG wrapper calls self.forward -> mask applied automatically.
half = x[: len(x) // 2]
combined = torch.cat([half, half], dim=0)
model_out = self.forward(combined, t, autoenc_cond, drop_mask)
eps, rest = model_out[:, : self.in_channels], model_out[:, self.in_channels:]
cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0)
half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps)
eps = torch.cat([half_eps, half_eps], dim=0)
return torch.cat([eps, rest], dim=1)
class SpatialDiffuseSlot(DiffuseSlot):
"""Semanticist DiffuseSlot with OUR spatial conditioning (85 multi-res tokens).
Extra params (config):
level_sizes : (8,4,2,1) -> 85 cond tokens
pool_depth : AttnPool depth (default 2, our L_repa choice)
freeze_dit : freeze pretrained DiT trunk (warmup phase 1)
NOTE: pass num_slots=256 in config so the tok_L ENCODER loads 100%; the DiT
is rebuilt with num_autoenc=85 (null_cond re-init is expected/normal).
"""
def __init__(self, *, level_sizes=(8, 4, 2, 1), pool_depth=2,
freeze_dit=False, init_from=None, dit_lr_scale=None,
dit_model="DiT-L-2", **kwargs):
super().__init__(dit_model=dit_model, **kwargs)
# Phase-2 knob: when the trunk is UNFROZEN, create_optimizer puts dit.*
# (except cond-embedder/null_cond) into `lr_scale`-scaled param groups —
# timm scheduler applies the scale to warmup AND cosine, so the
# pretrained trunk moves ~dit_lr_scale× slower (protects tok_L knowledge).
self.dit_lr_scale = dit_lr_scale
self.level_sizes = tuple(level_sizes)
self.num_cond_tokens = sum(s * s for s in self.level_sizes) # 85
# rebuild DiT with num_autoenc=85 + spatial mask (same size/class family)
import semanticist.stage1.diffuse_slot as ds
cfgmap = { # mirror DiT_with_autoenc_cond_*_2 constructors
"DiT-L-2": dict(depth=24, hidden_size=1024, patch_size=2, num_heads=16),
"DiT-XL-2": dict(depth=28, hidden_size=1152, patch_size=2, num_heads=16),
"DiT-B-2": dict(depth=12, hidden_size=768, patch_size=2, num_heads=12),
}[dit_model]
old = self.dit
self.dit = DiTSpatial(
input_size=self.dit_input_size,
in_channels=self.dit_in_channels,
num_autoenc=self.num_cond_tokens,
autoenc_dim=kwargs.get("slot_dim", 16),
use_repa=self.use_repa,
encoder_depth=old.encoder_depth if hasattr(old, "encoder_depth") else 8,
level_sizes=self.level_sizes,
**cfgmap,
)
del old
# our multi-res pool on top of the (frozen-init) Semanticist ViT patches
enc_d = self.encoder.embed_dim if hasattr(self.encoder, "embed_dim") else 768
self.spatial_pool = SpatialAttnPool(
enc_d=enc_d, level_sizes=self.level_sizes, depth=pool_depth)
# tok_L init (encoder 100% + DiT trunk; null_cond K-mismatch handled).
# NOTE: strict=False still ERRORS on shape mismatch -> pop mismatched keys.
if init_from:
# init_from can be:
# (a) tok_L .pkl -> Phase-1 pretrained init
# (b) an accelerate ckpt DIR (has model.safetensors) -> Phase-2:
# load the whole SpatialDiffuseSlot weights (encoder+pool+DiT+cond)
# from a prior frozen-run checkpoint, then continue with DiT unfrozen.
if os.path.isdir(init_from):
from safetensors.torch import load_file
sf = os.path.join(init_from, "model.safetensors")
ck = load_file(sf)
phase = "PHASE-2 (resume weights from ckpt dir)"
else:
ck = torch.load(init_from, map_location="cpu")
phase = "PHASE-1 (tok_L pretrained)"
ck = {k.replace("._orig_mod", ""): v for k, v in ck.items()}
own = self.state_dict()
dropped = [k for k, v in list(ck.items())
if k in own and own[k].shape != v.shape]
for k in dropped:
ck.pop(k)
ret = self.load_state_dict(ck, strict=False)
enc_miss = [k for k in ret.missing_keys if k.startswith("encoder.")]
dit_miss = [k for k in ret.missing_keys
if k.startswith("dit.") and "null_cond" not in k
and "pos_embed" not in k and "spatial_mask" not in k]
pool_miss = [k for k in ret.missing_keys if k.startswith("spatial_pool.")]
print(f"[SpatialDiffuseSlot] init_from={init_from} [{phase}]\n"
f" encoder missing={len(enc_miss)} dit-trunk missing={len(dit_miss)} "
f"pool missing={len(pool_miss)} shape-dropped={dropped}")
assert not enc_miss and not dit_miss, "init incomplete!"
# Phase-2 must also restore the trained attn-pool (Phase-1 doesn't have it)
if os.path.isdir(init_from):
assert not pool_miss, "Phase-2 resume: spatial_pool weights missing!"
# BUG#8 FIX: cond length is 85, but DiffuseSlot kept num_slots=256 (for the
# encoder ckpt). Trainer eval + nested_sampler size their drop_mask from
# num_slots -> (256) vs cond (85) crash at eval. Re-point BOTH to 85.
# (encoder keeps its own internal num_slots=256 attr -> ckpt still loads.)
self.num_slots = self.num_cond_tokens # 85 (trainer reads this)
# OUR level-drop (coarse-first keep, 8x8 dropped first, token-proportional).
self.nested_sampler = LevelNestedSampler(self.level_sizes)
self.freeze_dit = freeze_dit
if freeze_dit:
self._set_dit_trunk_grad(False)
def _set_dit_trunk_grad(self, flag: bool):
# trunk = pretrained parts (keep cond-embedder/null_cond trainable)
for name, p in self.dit.named_parameters():
if name.startswith(("autoenc_cond_embedder", "null_cond")):
p.requires_grad = True
else:
p.requires_grad = flag
# ---- encoder path: Semanticist ViT patches -> our pool -> 85 tokens ----
def encode_patches(self, x):
enc = self.encoder
h = enc.prepare_tokens(x)
# patches must not see slots (their causal mask already enforces this);
# replicate the mask so patch features match tok_L's training regime.
T = h.shape[1]
attn_mask = torch.ones(T, T, device=h.device, dtype=torch.bool)
ns = enc.num_slots
causal = torch.ones(ns, ns, device=h.device, dtype=torch.bool).tril(0)
attn_mask[-ns:, -ns:] = causal
attn_mask[:-ns, -ns:] = False
for blk in enc.blocks:
h = blk(h, attn_mask=attn_mask)
h = enc.norm(h)
num_patches = T - 1 - ns
return h[:, 1:1 + num_patches] # (B, 256, enc_d) drop cls/slots
def encode_slots(self, x):
feats = self.encode_patches(x)
tokens = self.spatial_pool(feats) # (B, 85, enc_d)
slots = self.encoder2slot(tokens) # reuse Linear(enc_d, slot_dim)
if self.norm_slots:
std = torch.std(slots, dim=-1, keepdim=True)
mean = torch.mean(slots, dim=-1, keepdim=True)
slots = (slots - mean) / (std + 1e-6)
return slots
|