"""Inference-only loader for SushiUI's merged style encoder artifact.""" from __future__ import annotations import json from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from safetensors.torch import load_file from transformers import Siglip2ImageProcessor from transformers.masking_utils import create_bidirectional_mask from transformers.models.siglip2.modeling_siglip2 import Siglip2VisionConfig, Siglip2VisionModel class AttentionPooling(nn.Module): def __init__(self, token_dim: int, heads: int) -> None: super().__init__() self.query = nn.Parameter(torch.zeros(1, 1, token_dim)) self.proj_k = nn.Linear(token_dim, token_dim) self.proj_v = nn.Linear(token_dim, token_dim) self.attn = nn.MultiheadAttention(token_dim, heads, batch_first=True) def forward(self, tokens: torch.Tensor, valid: torch.Tensor) -> torch.Tensor: query = self.query.expand(tokens.shape[0], -1, -1) pooled, _ = self.attn(query, self.proj_k(tokens), self.proj_v(tokens), key_padding_mask=~valid) return pooled.squeeze(1) class StyleEncoder(nn.Module): def __init__(self, architecture: dict) -> None: super().__init__() config = Siglip2VisionConfig(**{**architecture["vision_config"], "attn_implementation": "sdpa"}) self.vision_encoder = Siglip2VisionModel(config) self.feature_layers = tuple(architecture["feature_layers_resolved"]) self.depth = max(self.feature_layers) hidden = config.hidden_size token_dim = architecture["token_dim"] self.layer_norms = nn.ModuleList(nn.LayerNorm(hidden, eps=config.layer_norm_eps) for _ in self.feature_layers) self.token_proj = nn.Linear(hidden * len(self.feature_layers), token_dim) self.pooling = AttentionPooling(token_dim, architecture["pool_heads"]) self.projector = nn.Sequential(nn.Linear(token_dim, architecture["projector_hidden_dim"]), nn.GELU(), nn.Linear(architecture["projector_hidden_dim"], architecture["embed_dim"])) def forward(self, pixel_values: torch.Tensor, attention_mask: torch.Tensor, spatial_shapes: torch.Tensor) -> torch.Tensor: valid = attention_mask != 0 pixel_values = torch.where(valid.unsqueeze(-1), pixel_values, torch.zeros_like(pixel_values)) vision = self.vision_encoder hidden = vision.embeddings(pixel_values, spatial_shapes) mask = create_bidirectional_mask(config=vision.config, inputs_embeds=hidden, attention_mask=valid) states = [hidden] for layer in vision.encoder.layers[:self.depth]: hidden = layer(hidden, mask) states.append(hidden) tokens = torch.cat([norm(states[index]) for norm, index in zip(self.layer_norms, self.feature_layers)], dim=-1) tokens = self.token_proj(tokens) tokens = torch.where(valid.unsqueeze(-1), tokens, torch.zeros_like(tokens)) pooled = self.pooling(tokens, valid) return F.normalize(self.projector(pooled).float(), dim=-1) def load_encoder(encoder_path: str) -> tuple[StyleEncoder, Siglip2ImageProcessor, int]: path = Path(encoder_path) metadata = json.loads(path.with_name("best_metadata.json").read_text(encoding="utf-8")) if metadata.get("format") != "sushiui-style-encoder" or metadata.get("format_version") != 1: raise ValueError("Unsupported style encoder artifact") model = StyleEncoder(metadata["architecture"]) missing, unexpected = model.load_state_dict(load_file(str(path)), strict=False) unused = ("vision_encoder.post_layernorm.", "vision_encoder.head.") missing = [name for name in missing if not name.startswith(unused)] if missing or unexpected: raise ValueError(f"Artifact does not match architecture: {len(missing)} missing, " f"{len(unexpected)} unexpected tensors") processor = Siglip2ImageProcessor.from_dict(metadata["preprocessing"]) return model.eval(), processor, int(metadata["style_max_num_patches"])