Spaces:
Running on Zero
Running on Zero
Download model.py from cella110n/style_embedder: direct link, hf CLI and curl.
- Browser
- Download file 4.23 kB
-
https://huggingface.co/spaces/cella110n/style_embedder/resolve/main/model.py
- Command line
-
hf download hf://spaces/cella110n/style_embedder/model.py
-
curl -L -o model.py https://huggingface.co/spaces/cella110n/style_embedder/resolve/main/model.py
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"]) | |