a12s12's picture
Super-squash branch 'main' using huggingface_hub
f065e53
Raw
History Blame Contribute Delete
4.3 kB
"""Spatial-align attention mask for [latent | cond] concat self-attention.
Ports OUR cross-attn xa_mask (model_multires.py L1326: latent patch (r,c) on the
GxG latent grid attends, per level s, the cond token covering it) into a single
(N+M)x(N+M) self-attn mask for Semanticist's DiT (which concats cond to the
latent sequence and runs full self-attn).
Semantics preserved EXACTLY (2026-07-15 user requirement):
latent->latent : True (== our sa_mask=None, full)
latent->cond : our xa_mask spatial rule (same code as model_multires)
cond->latent : False (cond was never a query in cross-attn; conditioning
direction latent<-cond must be preserved)
cond->cond : identity only (attention must not mix cond tokens; every row
needs >=1 True or SDPA produces NaN)
Level-drop / uncond stays VALUE-based (learned null values), mask unchanged —
matches our learned_null_attend and Semanticist's null_cond.
"""
import torch
def build_xa_mask(G: int, level_sizes=(8, 4, 2, 1)) -> torch.Tensor:
"""(N_img, M_cond) bool — EXACT port of model_multires.py L1326 loop.
G = latent grid side (DiT-L patch2 on 16x16 latent -> G=8).
level_sizes: multi-res cond grids, coarse order irrelevant (offsets follow list order).
"""
num_img = G * G
offsets, off = {}, 0
for s in level_sizes:
offsets[s] = off
off += s * s
M = off
xa = torch.zeros(num_img, M, dtype=torch.bool)
for s in level_sizes:
start = offsets[s]
for r in range(G):
for c in range(G):
img_idx = r * G + c
if G >= s:
xa[img_idx, start + (r * s // G) * s + (c * s // G)] = True
else:
ratio = s // G
for dr in range(ratio):
for dc in range(ratio):
xa[img_idx, start + (r * ratio + dr) * s + (c * ratio + dc)] = True
return xa
def build_concat_self_attn_mask(G: int, level_sizes=(8, 4, 2, 1)) -> torch.Tensor:
"""(N+M, N+M) bool self-attn mask, True=attend. Sequence = [latent N | cond M]."""
xa = build_xa_mask(G, level_sizes)
N, M = xa.shape
T = N + M
m = torch.zeros(T, T, dtype=torch.bool)
m[:N, :N] = True # latent -> latent : full
m[:N, N:] = xa # latent -> cond : spatial routing
# cond -> latent : False (stays zero)
m[N:, N:] = torch.eye(M, dtype=torch.bool) # cond -> cond : identity
return m
if __name__ == "__main__":
# ---- unit tests (CPU) ----
G, LS = 8, (8, 4, 2, 1)
xa = build_xa_mask(G, LS)
assert xa.shape == (64, 85), xa.shape
# per latent patch: exactly one token per level -> 4 True per row
assert (xa.sum(1) == len(LS)).all(), xa.sum(1)
# spot checks: latent (0,0)=idx0 -> lvl8 tok0, lvl4 tok0(off64), lvl2 tok0(off80), lvl1 tok0(off84)
assert xa[0, 0] and xa[0, 64] and xa[0, 80] and xa[0, 84]
# latent (7,7)=idx63 -> lvl8 tok63, lvl4 (3,3)=off64+15, lvl2 (1,1)=off80+3, lvl1 off84
assert xa[63, 63] and xa[63, 64 + 15] and xa[63, 80 + 3] and xa[63, 84]
# latent (3,4)=idx28 -> lvl8 tok28, lvl4 (1,2)=off64+6, lvl2 (0,1)=off80+1, lvl1
assert xa[28, 28] and xa[28, 64 + 6] and xa[28, 80 + 1] and xa[28, 84]
# no cross-cell leakage: latent 0 must NOT see lvl8 tok1 nor lvl4 tok5
assert not xa[0, 1] and not xa[0, 64 + 5]
m = build_concat_self_attn_mask(G, LS)
T = 64 + 85
assert m.shape == (T, T)
assert m[:64, :64].all() # latent full
assert not m[64:, :64].any() # cond -> latent blocked
assert (m[64:, 64:] == torch.eye(85, dtype=torch.bool)).all() # cond identity
assert (m.sum(1) >= 1).all() # no all-False row (SDPA NaN guard)
# verify latent->cond block equals xa exactly
assert (m[:64, 64:] == xa).all()
# SDPA smoke: mask must produce finite output
import torch.nn.functional as F
q = torch.randn(2, 4, T, 32); k = torch.randn(2, 4, T, 32); v = torch.randn(2, 4, T, 32)
out = F.scaled_dot_product_attention(q, k, v, attn_mask=m)
assert torch.isfinite(out).all()
print("ALL MASK TESTS PASSED ✅ (xa 64x85, concat 149x149, SDPA finite)")