prism / layers.py
litcoderr's picture
Publish PRISM weights and modeling code
a596b0a verified
Raw
History Blame
4.17 kB
"""Shared building blocks for PRISM.
- ``QFormerBlock`` : one block of the per-frame Q-Former (self-attn over the
queries → cross-attn into the frozen patch tokens → FFN). Used by the
Decompositional Encoder's frame-level stage.
- ``TemporalBlock`` : a pre-LN transformer block used by both the temporal
stage of the Decompositional Encoder (causal) and the Compositional
Latent Predictor.
- ``build_sin_pos_embed`` / ``causal_mask`` : positional-embedding and masking
helpers.
PyTorch's ``MultiheadAttention`` boolean masks follow the convention
``True == blocked`` for both ``attn_mask`` and ``key_padding_mask``.
"""
from __future__ import annotations
import math
import torch
from torch import nn
def build_sin_pos_embed(num_positions: int, dim: int) -> torch.Tensor:
"""1-D sinusoidal positional embedding, shape ``(1, num_positions, dim)``."""
pe = torch.zeros(num_positions, dim)
pos = torch.arange(0, num_positions, dtype=torch.float32).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, dim, 2, dtype=torch.float32) * (-math.log(10000.0) / dim)
)
pe[:, 0::2] = torch.sin(pos * div_term)
pe[:, 1::2] = torch.cos(pos * div_term)
return pe.unsqueeze(0)
def causal_mask(length: int, device: torch.device | None = None) -> torch.Tensor:
"""Boolean causal mask ``(L, L)``; ``True`` blocks future positions."""
return torch.triu(
torch.ones(length, length, dtype=torch.bool, device=device), diagonal=1
)
class QFormerBlock(nn.Module):
"""BLIP-2-style block: query self-attn → query→patch cross-attn → FFN."""
def __init__(
self,
d_z: int,
d_kv: int,
num_heads: int = 8,
mlp_ratio: float = 4.0,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
):
super().__init__()
self.norm1 = nn.LayerNorm(d_z)
self.self_attn = nn.MultiheadAttention(
embed_dim=d_z, num_heads=num_heads, dropout=attn_drop, batch_first=True
)
self.norm2_q = nn.LayerNorm(d_z)
self.norm2_kv = nn.LayerNorm(d_kv)
self.cross_attn = nn.MultiheadAttention(
embed_dim=d_z,
num_heads=num_heads,
dropout=attn_drop,
kdim=d_kv,
vdim=d_kv,
batch_first=True,
)
self.norm3 = nn.LayerNorm(d_z)
hidden = int(d_z * mlp_ratio)
self.mlp = nn.Sequential(
nn.Linear(d_z, hidden),
nn.GELU(),
nn.Dropout(proj_drop),
nn.Linear(hidden, d_z),
nn.Dropout(proj_drop),
)
def forward(self, q: torch.Tensor, kv: torch.Tensor) -> torch.Tensor:
h = self.norm1(q)
sa, _ = self.self_attn(h, h, h, need_weights=False)
q = q + sa
hq = self.norm2_q(q)
hkv = self.norm2_kv(kv)
ca, _ = self.cross_attn(hq, hkv, hkv, need_weights=False)
q = q + ca
q = q + self.mlp(self.norm3(q))
return q
class TemporalBlock(nn.Module):
"""Pre-LN transformer block (self-attn + FFN).
Accepts a boolean ``attn_mask`` ``(L, L)`` (True = blocked) and a
``key_padding_mask`` ``(B, L)`` (True = padded) on every forward.
"""
def __init__(self, d_in: int, num_heads: int = 8, mlp_ratio: float = 4.0):
super().__init__()
self.norm1 = nn.LayerNorm(d_in)
self.attn = nn.MultiheadAttention(
embed_dim=d_in, num_heads=num_heads, batch_first=True
)
self.norm2 = nn.LayerNorm(d_in)
hidden = int(d_in * mlp_ratio)
self.mlp = nn.Sequential(
nn.Linear(d_in, hidden),
nn.GELU(),
nn.Linear(hidden, d_in),
)
def forward(
self,
x: torch.Tensor,
attn_mask: torch.Tensor | None = None,
key_padding_mask: torch.Tensor | None = None,
) -> torch.Tensor:
h = self.norm1(x)
a, _ = self.attn(
h, h, h,
need_weights=False,
attn_mask=attn_mask,
key_padding_mask=key_padding_mask,
)
x = x + a
x = x + self.mlp(self.norm2(x))
return x