style_embedder / model.py
cella110n's picture
Add gated style A/B comparison UI
20b2a0b verified
Raw History Blame Contribute Delete
4.23 kB
"""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"])