OmniRet / omniret /modules.py
chuonghm's picture
Publish OmniRet checkpoint-2220
03f1aba verified
Raw
History Blame Contribute Delete
10 kB
from __future__ import annotations
import torch
from torch import nn
import torch.nn.functional as F
class MLPProjector(nn.Module):
def __init__(self, in_dim: int, out_dim: int) -> None:
super().__init__()
self.net = nn.Sequential(nn.Linear(in_dim, out_dim), nn.GELU(), nn.Linear(out_dim, out_dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.is_inference():
x = x.clone()
return self.net(x.to(dtype=self.net[0].weight.dtype))
def _feed_forward(dim: int, mult: int = 1) -> nn.Sequential:
# Matches the reference `FeedForward` (perceiver.py / embedding_proj.py):
# pre-LayerNorm, bias-free linears, ff_mult=1.
inner = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner, bias=False),
nn.GELU(),
nn.Linear(inner, dim, bias=False),
)
class PerceiverAttention(nn.Module):
"""Reference `PerceiverAttention` block (perceiver.py:103).
Latents attend to the concatenation of themselves and the media/token
context. Dual pre-norm (media + latents), bias-free q/kv/out projections.
"""
def __init__(self, dim: int, dim_head: int = 64, heads: int = 8) -> None:
super().__init__()
self.heads = heads
self.dim_head = dim_head
inner_dim = dim_head * heads
self.norm_media = nn.LayerNorm(dim)
self.norm_latents = nn.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
b, n, _ = x.shape
return x.view(b, n, self.heads, self.dim_head).transpose(1, 2) # (B, H, N, d)
def forward(self, x: torch.Tensor, latents: torch.Tensor, key_padding_mask: torch.Tensor | None = None) -> torch.Tensor:
x = self.norm_media(x)
latents = self.norm_latents(latents)
b, num_latents, _ = latents.shape
q = self._split_heads(self.to_q(latents))
kv_input = torch.cat([latents, x], dim=1)
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
k = self._split_heads(k)
v = self._split_heads(v)
attn_mask = None
if key_padding_mask is not None:
# key_padding_mask: (B, N) True = valid media token. Latents are
# always attendable. SDPA bool mask: True = participate in attention.
latent_valid = torch.ones(b, num_latents, dtype=torch.bool, device=key_padding_mask.device)
valid = torch.cat([latent_valid, key_padding_mask.to(torch.bool)], dim=1) # (B, L+N)
attn_mask = valid[:, None, None, :] # (B, 1, 1, L+N), broadcast over heads/queries
out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
out = out.transpose(1, 2).reshape(b, num_latents, self.heads * self.dim_head)
return self.to_out(out)
def _kaiming_init(module: nn.Module) -> None:
# Reference init (perceiver.py:291, embedding_proj.py:350): kaiming_uniform_
# on linears, LayerNorm weight=1 / bias=0.
if isinstance(module, nn.Linear):
nn.init.kaiming_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
elif isinstance(module, nn.LayerNorm):
nn.init.constant_(module.bias, 0)
nn.init.constant_(module.weight, 1.0)
class SharedMediaResampler(nn.Module):
"""Media resampler matching the reference `perceiver_v2` block structure.
Reference `PerceiverResamplerV2` (perceiver.py:261): depth-2 PerceiverAttention
blocks (dual-norm, bias-free, 8 heads, dim_head = dim // 8), bias-free ff_mult=1
FeedForward, a final LayerNorm, kaiming init, and additive per-modality latent
embeddings. We keep the `shared + modality` latent split so Stage-2 can inherit
the image latents for video (`initialize_stage2_from_state_dict`).
"""
def __init__(self, dim: int, num_latents: int = 64, num_heads: int = 8, depth: int = 2) -> None:
super().__init__()
self.shared_latents = nn.Parameter(torch.randn(num_latents, dim))
self.modality_latents = nn.ParameterDict(
{name: nn.Parameter(torch.randn(num_latents, dim)) for name in ("image", "video", "audio")}
)
dim_head = max(1, dim // num_heads)
self.layers = nn.ModuleList(
nn.ModuleList([PerceiverAttention(dim, dim_head=dim_head, heads=num_heads), _feed_forward(dim, mult=1)])
for _ in range(depth)
)
self.norm = nn.LayerNorm(dim)
self.apply(_kaiming_init)
def forward(self, tokens: torch.Tensor, modality: str, mask: torch.Tensor | None = None) -> torch.Tensor:
if modality not in self.modality_latents:
raise ValueError(f"unknown modality: {modality}")
latents = self.shared_latents + self.modality_latents[modality]
latents = latents.unsqueeze(0).expand(tokens.shape[0], -1, -1)
key_padding_mask = None if mask is None else mask.to(device=tokens.device, dtype=torch.bool)
for attn, ff in self.layers:
latents = attn(tokens, latents, key_padding_mask) + latents
latents = ff(latents) + latents
return self.norm(latents)
class PerceiverLatentAttention(nn.Module):
"""In-pooler latent-attention pre-pooler used by the reference `sw_attention`
head (`SWProj(using_latent_attention=True)`, embedding_proj.py:342).
Compresses a variable-length (masked) token sequence into `num_latents` dense
latent tokens via a depth-2 PerceiverResampler (dim_head 256, 4 heads, ff_mult 1,
final LayerNorm). This is the content-adaptive component that lets the pooler
down-weight the shared instruction/BOS tokens; without it the sliced-Wasserstein
projection pools the raw frozen sequence and embeddings cannot separate.
"""
def __init__(self, dim: int, num_latents: int, depth: int = 2, dim_head: int = 256, heads: int = 4, ff_mult: int = 1) -> None:
super().__init__()
self.latents = nn.Parameter(torch.randn(num_latents, dim))
self.layers = nn.ModuleList(
nn.ModuleList([PerceiverAttention(dim, dim_head=dim_head, heads=heads), _feed_forward(dim, mult=ff_mult)])
for _ in range(depth)
)
self.norm = nn.LayerNorm(dim)
self.apply(_kaiming_init)
def forward(self, tokens: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:
latents = self.latents.unsqueeze(0).expand(tokens.shape[0], -1, -1)
key_padding_mask = None if mask is None else mask.to(device=tokens.device, dtype=torch.bool)
for attn, ff in self.layers:
latents = attn(tokens, latents, key_padding_mask) + latents
latents = ff(latents) + latents
return self.norm(latents) # (B, num_latents, dim), no padding
class AttentionSlicedWassersteinPooling(nn.Module):
"""Reference `sw_attention` pooler = latent-attention pre-pooler + sliced-Wasserstein
embedding (`RandPSWE`, reference_init='uniform', pooling_method='max').
The token sequence is first compressed to `num_references` dense latents, then the
sliced-Wasserstein step sorts the latent projections against a learnable uniform
reference and pools with a straight-through max. Because the latents are dense and
already count `num_references`, the SW step needs no mask or grid-sample
interpolation (reference runs `sample=False`).
"""
def __init__(
self,
dim: int,
num_references: int = 128,
num_projections: int = 4096,
depth: int = 2,
dim_head: int = 256,
heads: int = 4,
) -> None:
super().__init__()
self.num_references = num_references
self.latent_attention = PerceiverLatentAttention(dim, num_references, depth=depth, dim_head=dim_head, heads=heads)
theta_v = torch.normal(mean=0, std=1, size=(dim, num_projections))
if num_projections <= dim:
theta_v = torch.eye(dim, num_projections, dtype=torch.float32).type(theta_v.dtype)
self.theta_v = nn.Parameter(theta_v)
# Learnable uniform reference, matching reference RandPSWE reference_init='uniform'
# (embedding_proj.py:92-94): a full (num_references, num_projections) parameter
# initialised to a linspace ramp, not the affine base_reference variant.
uniform_ref = torch.linspace(-1, 1, num_references).unsqueeze(1).repeat(1, num_projections)
self.reference = nn.Parameter(uniform_ref)
@property
def theta(self) -> torch.Tensor:
return self.theta_v
def get_slice(self, X: torch.Tensor) -> torch.Tensor:
theta_v = torch.nan_to_num(self.theta_v, nan=0.0, posinf=0.0, neginf=0.0)
theta_norm = torch.linalg.vector_norm(theta_v, dim=0, keepdim=True)
theta_norm = torch.nan_to_num(theta_norm, nan=0.0, posinf=0.0, neginf=0.0).clamp(min=1e-8)
theta = theta_v / theta_norm
theta = torch.nan_to_num(theta, nan=0.0, posinf=0.0, neginf=0.0)
return torch.matmul(X, theta.to(X.dtype).to(X.device))
def forward(self, tokens: torch.Tensor, mask: torch.Tensor | None = None) -> torch.Tensor:
latents = self.latent_attention(tokens, mask) # (B, R, D)
slices = self.get_slice(latents) # (B, R, num_projections)
sorted_slices = torch.sort(slices, dim=1).values
refs = self.reference.to(slices.dtype).unsqueeze(0).expand_as(sorted_slices)
_, ref_order = torch.sort(refs, dim=1)
sorted_slices = torch.gather(sorted_slices, dim=1, index=ref_order)
coupled = refs - sorted_slices
if not self.training:
return coupled.max(dim=1).values
y_soft = F.softmax(coupled, dim=1)
indices = torch.argmax(y_soft, dim=1, keepdim=True)
y_hard = torch.zeros_like(coupled).scatter_(1, indices, 1)
return (coupled * (y_hard - y_soft.detach() + y_soft)).sum(dim=1)