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