SolPix / solpix /model.py
j0no12's picture
Publish SolPix step 210000 research checkpoint
8d31176 verified
Raw History Blame Contribute Delete
11.8 kB
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
from torch import Tensor, nn
import torch.nn.functional as F
from .config import SolPixConfig
@dataclass
class SolPixOutput:
velocity: Tensor
# Present only when return_features=True; shaped [batch, image_tokens, width].
repa_features: Tensor | None = None
def _sinusoidal_embedding(values: Tensor, width: int, max_period: float = 10_000.0) -> Tensor:
"""Sinusoidal embedding for scalar values, with a final dimension of width."""
half = width // 2
frequencies = torch.exp(
-math.log(max_period)
* torch.arange(half, device=values.device, dtype=torch.float32)
/ max(half - 1, 1)
)
phases = values.float().reshape(-1, 1) * frequencies.reshape(1, -1)
embedding = torch.cat((phases.cos(), phases.sin()), dim=-1)
if width % 2:
embedding = F.pad(embedding, (0, 1))
return embedding
def _image_position_embedding(height: int, width: int, dim: int, device: torch.device) -> Tensor:
"""Parameter-free 2D sine/cosine positions for a variable latent grid."""
quarter = dim // 4
frequencies = torch.exp(
-math.log(10_000.0)
* torch.arange(quarter, device=device, dtype=torch.float32)
/ max(quarter - 1, 1)
)
y = (torch.arange(height, device=device, dtype=torch.float32) + 0.5) / height * 2 - 1
x = (torch.arange(width, device=device, dtype=torch.float32) + 0.5) / width * 2 - 1
yy, xx = torch.meshgrid(y, x, indexing="ij")
x_phase = xx.reshape(-1, 1) * frequencies.reshape(1, -1) * math.pi
y_phase = yy.reshape(-1, 1) * frequencies.reshape(1, -1) * math.pi
return torch.cat((x_phase.sin(), x_phase.cos(), y_phase.sin(), y_phase.cos()), dim=-1)
class TimestepConditioner(nn.Module):
"""Build one shared adaLN conditioning vector from time and image geometry."""
def __init__(self, width: int):
super().__init__()
self.time_mlp = nn.Sequential(nn.Linear(width, width), nn.SiLU(), nn.Linear(width, width))
self.geometry_mlp = nn.Sequential(nn.Linear(4, width), nn.SiLU(), nn.Linear(width, width))
self.modulation = nn.Linear(width, 6 * width)
# adaLN-Zero style residual gates: start near the identity function.
nn.init.zeros_(self.modulation.weight)
nn.init.zeros_(self.modulation.bias)
def forward(self, time: Tensor, latent_height: int, latent_width: int, factor: int) -> Tensor:
width_px = latent_width * factor
height_px = latent_height * factor
geometry = torch.tensor(
[
math.log(max(height_px, 1) / 512.0),
math.log(max(width_px, 1) / 512.0),
math.log(max(width_px, 1) / max(height_px, 1)),
math.log(max(height_px, width_px, 1) / 512.0),
],
device=time.device,
dtype=torch.float32,
).expand(time.shape[0], -1)
t_embed = _sinusoidal_embedding(time, self.time_mlp[0].in_features)
hidden = self.time_mlp(t_embed) + self.geometry_mlp(geometry)
# [B, 6, D]: attention shift/scale/gate, then MLP shift/scale/gate.
return self.modulation(hidden).reshape(time.shape[0], 6, -1)
class SolPixBlock(nn.Module):
def __init__(self, config: SolPixConfig):
super().__init__()
dim = config.width
self.heads = config.heads
self.head_dim = dim // config.heads
self.norm_attn = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
self.qkv = nn.Linear(dim, 3 * dim)
self.attn_out = nn.Linear(dim, dim)
self.norm_mlp = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
self.mlp_in = nn.Linear(dim, 2 * config.mlp_hidden)
self.local_mix = nn.Conv2d(
config.mlp_hidden,
config.mlp_hidden,
kernel_size=3,
padding=1,
groups=config.mlp_hidden,
)
self.mlp_out = nn.Linear(config.mlp_hidden, dim)
@staticmethod
def _adaptive_norm(x: Tensor, norm: nn.LayerNorm, shift: Tensor, scale: Tensor) -> Tensor:
return norm(x) * (1 + scale[:, None, :]) + shift[:, None, :]
def forward(
self,
tokens: Tensor,
condition: Tensor,
text_mask: Tensor,
text_count: int,
latent_height: int,
latent_width: int,
) -> Tensor:
shift_a, scale_a, gate_a, shift_m, scale_m, gate_m = condition.unbind(dim=1)
batch, token_count, dim = tokens.shape
normalized = self._adaptive_norm(tokens, self.norm_attn, shift_a, scale_a)
qkv = self.qkv(normalized).reshape(batch, token_count, 3, self.heads, self.head_dim)
q, k, v = qkv.permute(2, 0, 3, 1, 4).unbind(dim=0)
# SDPA boolean masks use True for keys that may participate in attention.
image_keys = torch.ones(
(batch, token_count - text_count), device=tokens.device, dtype=torch.bool
)
key_mask = torch.cat((text_mask, image_keys), dim=1)[:, None, None, :]
attended = F.scaled_dot_product_attention(q, k, v, attn_mask=key_mask)
attended = attended.transpose(1, 2).reshape(batch, token_count, dim)
tokens = tokens + gate_a[:, None, :] * self.attn_out(attended)
normalized = self._adaptive_norm(tokens, self.norm_mlp, shift_m, scale_m)
gate, value = self.mlp_in(normalized).chunk(2, dim=-1)
image_value = value[:, text_count:, :].transpose(1, 2).reshape(
batch, -1, latent_height, latent_width
)
image_value = image_value + self.local_mix(image_value)
image_value = image_value.flatten(2).transpose(1, 2)
value = torch.cat((value[:, :text_count, :], image_value), dim=1)
tokens = tokens + gate_m[:, None, :] * self.mlp_out(F.silu(gate) * value)
return tokens
class SolPixTransformer2D(nn.Module):
"""U-shaped joint text/image flow transformer operating on frozen AE latents.
Inputs are precomputed text-encoder states and noisy autoencoder latents. The
external text encoder and autoencoder are intentionally not registered here.
"""
def __init__(self, config: SolPixConfig | None = None):
super().__init__()
self.config = config or SolPixConfig()
c = self.config
self.text_projection = nn.Linear(c.text_dim, c.width)
self.latent_projection = nn.Linear(c.latent_channels, c.width)
self.modality_embedding = nn.Parameter(torch.empty(2, c.width))
nn.init.normal_(self.modality_embedding, std=0.02)
self.conditioner = TimestepConditioner(c.width)
self.blocks = nn.ModuleList(SolPixBlock(c) for _ in range(c.depth))
self.skip_projections = nn.ModuleList(
nn.Linear(2 * c.width, c.width) for _ in range(c.encoder_depth)
)
self.final_norm = nn.LayerNorm(c.width, elementwise_affine=False, eps=1e-6)
self.latent_output = nn.Linear(c.width, c.latent_channels)
self.output_mix = nn.Conv2d(c.latent_channels, c.latent_channels, kernel_size=3, padding=1)
def forward(
self,
noisy_latents: Tensor,
time: Tensor,
text_embeddings: Tensor,
text_mask: Tensor | None = None,
*,
return_features: bool = False,
) -> Tensor | SolPixOutput:
"""Predict flow velocity for latent tensors shaped [B,C,H,W].
``text_mask`` is True for real text tokens and False for padding. To train
classifier-free guidance, encode empty prompts with the external text
encoder and pass those states for the selected examples.
"""
if noisy_latents.ndim != 4:
raise ValueError("noisy_latents must have shape [batch, channels, height, width]")
batch, channels, height, width = noisy_latents.shape
if channels != self.config.latent_channels:
raise ValueError(f"expected {self.config.latent_channels} latent channels, got {channels}")
if text_embeddings.ndim != 3 or text_embeddings.shape[0] != batch:
raise ValueError("text_embeddings must have shape [batch, text_tokens, text_dim]")
if text_embeddings.shape[1] > self.config.max_text_tokens:
raise ValueError(f"text sequence exceeds max_text_tokens={self.config.max_text_tokens}")
if text_embeddings.shape[2] != self.config.text_dim:
raise ValueError(f"expected text_dim={self.config.text_dim}, got {text_embeddings.shape[2]}")
if time.ndim == 0:
time = time.expand(batch)
if time.shape != (batch,):
raise ValueError("time must be a scalar or a [batch] tensor")
if text_mask is None:
text_mask = torch.ones(text_embeddings.shape[:2], device=text_embeddings.device, dtype=torch.bool)
elif text_mask.shape != text_embeddings.shape[:2]:
raise ValueError("text_mask must match the first two text_embeddings dimensions")
else:
text_mask = text_mask.to(device=text_embeddings.device, dtype=torch.bool)
text_count = text_embeddings.shape[1]
text = self.text_projection(text_embeddings)
modality_embedding = self.modality_embedding.to(dtype=text.dtype)
text = text + modality_embedding[0]
pos = _image_position_embedding(height, width, self.config.width, noisy_latents.device)
image = self.latent_projection(noisy_latents.permute(0, 2, 3, 1).reshape(batch, height * width, channels))
image = image + pos.to(dtype=image.dtype)[None, :, :] + modality_embedding[1]
tokens = torch.cat((text, image), dim=1)
condition = self.conditioner(
time, height, width, self.config.latent_downsample_factor
).to(dtype=tokens.dtype)
skips: list[Tensor] = []
repa_features: Tensor | None = None
for index in range(self.config.encoder_depth):
tokens = self.blocks[index](tokens, condition, text_mask, text_count, height, width)
skips.append(tokens)
if index == self.config.repa_layer and return_features:
repa_features = tokens[:, text_count:, :]
bottleneck_index = self.config.encoder_depth
tokens = self.blocks[bottleneck_index](tokens, condition, text_mask, text_count, height, width)
if bottleneck_index == self.config.repa_layer and return_features:
repa_features = tokens[:, text_count:, :]
for decoder_offset, (projection, skip) in enumerate(
zip(self.skip_projections, reversed(skips))
):
tokens = projection(torch.cat((tokens, skip), dim=-1))
block_index = bottleneck_index + 1 + decoder_offset
tokens = self.blocks[block_index](tokens, condition, text_mask, text_count, height, width)
if block_index == self.config.repa_layer and return_features:
repa_features = tokens[:, text_count:, :]
image = tokens[:, text_count:, :]
# Reuse the shared adaLN single modulation for the final adaptive norm.
shift, scale = condition[:, 0, :], condition[:, 1, :]
image = self.final_norm(image) * (1 + scale[:, None, :]) + shift[:, None, :]
velocity = self.latent_output(image).reshape(batch, height, width, channels).permute(0, 3, 1, 2)
velocity = self.output_mix(velocity)
if return_features:
return SolPixOutput(velocity=velocity, repa_features=repa_features)
return velocity
def parameter_count(self) -> int:
"""Count registered parameters, including any optional training heads."""
return sum(parameter.numel() for parameter in self.parameters())
# Keep the original training API working while exposing the model-card class name.
SolPix = SolPixTransformer2D