from __future__ import annotations import importlib import sys from pathlib import Path from typing import Dict, List, Optional import torch import torch.nn as nn import torch.nn.functional as F def _masked_softmax(logits: torch.Tensor, mask: torch.Tensor, dim: int = -1) -> torch.Tensor: mask = mask.to(dtype=torch.bool) masked_logits = logits.float().masked_fill(~mask, torch.finfo(torch.float32).min) weights = torch.softmax(masked_logits, dim=dim) weights = weights * mask.to(dtype=weights.dtype) normalizer = weights.sum(dim=dim, keepdim=True).clamp_min(1e-6) return (weights / normalizer).to(dtype=logits.dtype) def _stochastic_keep_mask(mask: torch.Tensor, drop_prob: float, training: bool) -> torch.Tensor: if (not training) or drop_prob <= 0.0: return mask keep_mask = mask.to(dtype=torch.bool).clone() batch_size, width = keep_mask.shape random_values = torch.rand(batch_size, width, device=mask.device) for idx in range(batch_size): available = torch.nonzero(keep_mask[idx], as_tuple=False).flatten() if available.numel() <= 1: continue keep = random_values[idx, available] > drop_prob if not keep.any(): keep[random_values[idx, available].argmax()] = True keep_mask[idx].fill_(False) keep_mask[idx, available] = keep return keep_mask.to(dtype=mask.dtype) def _grayscale(x: torch.Tensor) -> torch.Tensor: return 0.299 * x[:, 0:1] + 0.587 * x[:, 1:2] + 0.114 * x[:, 2:3] def _blur(x: torch.Tensor, kernel_size: int) -> torch.Tensor: if kernel_size <= 1: return x pad = kernel_size // 2 padded = F.pad(x, (pad, pad, pad, pad), mode="reflect") return F.avg_pool2d(padded, kernel_size=kernel_size, stride=1) def _sobel_edges(x: torch.Tensor) -> torch.Tensor: gray = _grayscale(x) kernel_x = torch.tensor( [[[-1.0, 0.0, 1.0], [-2.0, 0.0, 2.0], [-1.0, 0.0, 1.0]]], device=x.device, dtype=x.dtype, ).unsqueeze(0) kernel_y = torch.tensor( [[[-1.0, -2.0, -1.0], [0.0, 0.0, 0.0], [1.0, 2.0, 1.0]]], device=x.device, dtype=x.dtype, ).unsqueeze(0) gray = F.pad(gray, (1, 1, 1, 1), mode="reflect") grad_x = F.conv2d(gray, kernel_x) grad_y = F.conv2d(gray, kernel_y) magnitude = torch.sqrt(grad_x.pow(2) + grad_y.pow(2) + 1e-6) magnitude = magnitude / magnitude.amax(dim=(-2, -1), keepdim=True).clamp_min(1e-6) return magnitude.repeat(1, 3, 1, 1) def _color_tone(x: torch.Tensor) -> torch.Tensor: low_freq = _blur(x, kernel_size=11) luminance = 0.299 * low_freq[:, 0:1] + 0.587 * low_freq[:, 1:2] + 0.114 * low_freq[:, 2:3] cb = torch.clamp(low_freq[:, 2:3] - luminance + 0.5, 0.0, 1.0) cr = torch.clamp(low_freq[:, 0:1] - luminance + 0.5, 0.0, 1.0) return torch.cat([torch.clamp(luminance, 0.0, 1.0), cb, cr], dim=1) class LowLevelCueEncoder(nn.Module): def __init__(self, output_dim: int) -> None: super().__init__() self.net = nn.Sequential( nn.Conv2d(3, 24, kernel_size=3, stride=2, padding=1, bias=False), nn.BatchNorm2d(24), nn.SiLU(inplace=True), nn.Conv2d(24, 48, kernel_size=3, stride=2, padding=1, bias=False), nn.BatchNorm2d(48), nn.SiLU(inplace=True), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(48, output_dim), nn.LayerNorm(output_dim), ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.net(x) class ArcFaceHead(nn.Module): def __init__(self, num_classes: int, embed_dim: int, scale: float = 30.0) -> None: super().__init__() self.scale = float(scale) self.weight = nn.Parameter(torch.randn(num_classes, embed_dim)) nn.init.normal_(self.weight, std=0.02) def forward(self, embedding: torch.Tensor) -> Dict[str, torch.Tensor]: normalized_embedding = F.normalize(embedding, dim=-1) normalized_weight = F.normalize(self.weight, dim=-1) cosine = F.linear(normalized_embedding, normalized_weight) logits = cosine * self.scale return { "arcface_cosine": cosine, "class_logits": logits, } class DinoV3Backbone(nn.Module): def __init__( self, repo_dir: str, entrypoint: str, weights_path: str, freeze_backbone: bool = True, unfreeze_last_n_blocks: int = 0, ) -> None: super().__init__() self.repo_dir = str(Path(repo_dir).resolve()) self.entrypoint = entrypoint self.weights_path = str(Path(weights_path).resolve()) self.freeze_backbone = bool(freeze_backbone) self.unfreeze_last_n_blocks = max(0, int(unfreeze_last_n_blocks)) if self.repo_dir not in sys.path: sys.path.insert(0, self.repo_dir) backbone_module = importlib.import_module("dinov3.hub.backbones") backbone_fn = getattr(backbone_module, self.entrypoint) self.backbone = backbone_fn(weights=self.weights_path) self.hidden_size = int(getattr(self.backbone, "embed_dim")) self.num_register_tokens = int(getattr(self.backbone, "n_storage_tokens", 0) or 0) self.patch_size = int(getattr(self.backbone, "patch_size", 16)) self.register_buffer("pixel_mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1), persistent=False) self.register_buffer("pixel_std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1), persistent=False) self._configure_trainability() def _iter_encoder_blocks(self) -> List[nn.Module]: blocks = getattr(self.backbone, "blocks", None) if isinstance(blocks, (nn.ModuleList, list, tuple)): return list(blocks) return [] def _configure_trainability(self) -> None: for parameter in self.backbone.parameters(): parameter.requires_grad = not self.freeze_backbone if self.freeze_backbone and self.unfreeze_last_n_blocks > 0: blocks = self._iter_encoder_blocks() for parameter in self.backbone.parameters(): parameter.requires_grad = False for block in blocks[-self.unfreeze_last_n_blocks :]: for parameter in block.parameters(): parameter.requires_grad = True def train(self, mode: bool = True) -> "DinoV3Backbone": super().train(mode) all_frozen = not any(parameter.requires_grad for parameter in self.backbone.parameters()) if all_frozen: self.backbone.eval() else: self.backbone.train(mode) return self def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]: normalized = (x - self.pixel_mean.to(dtype=x.dtype, device=x.device)) / self.pixel_std.to(dtype=x.dtype, device=x.device) outputs = self.backbone.forward_features(normalized) pooled = outputs["x_norm_clstoken"] patch_tokens = outputs["x_norm_patchtokens"] storage_tokens = outputs.get("x_storage_tokens") if storage_tokens is None: storage_tokens = pooled.new_empty((pooled.size(0), 0, pooled.size(-1))) patch_hw = (normalized.size(-2) // self.patch_size, normalized.size(-1) // self.patch_size) tokens = torch.cat([pooled.unsqueeze(1), storage_tokens, patch_tokens], dim=1) return { "cls": pooled, "patch_tokens": patch_tokens, "storage_tokens": storage_tokens, "patch_hw": patch_hw, "token_sequence": tokens, } class SharedDinoViewBranch(nn.Module): def __init__(self, branch_type: str, token_dim: int, hidden_dim: int, output_dim: int) -> None: super().__init__() self.branch_type = branch_type self.view_embeddings = nn.Embedding(3, hidden_dim) self.missing_embeddings = nn.Embedding(3, hidden_dim) self.token_prompt = nn.Parameter(torch.randn(token_dim)) self.token_proj = nn.Sequential( nn.LayerNorm(token_dim * 2), nn.Linear(token_dim * 2, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, hidden_dim), nn.LayerNorm(hidden_dim), ) self.lowlevel_encoder = LowLevelCueEncoder(hidden_dim) self.merge = nn.Sequential( nn.LayerNorm(hidden_dim * 2), nn.Linear(hidden_dim * 2, hidden_dim), nn.GELU(), nn.Dropout(0.1), nn.Linear(hidden_dim, hidden_dim), nn.LayerNorm(hidden_dim), ) self.view_gate = nn.Sequential( nn.Linear(hidden_dim, hidden_dim), nn.SiLU(inplace=True), nn.Linear(hidden_dim, 1), ) self.out_proj = nn.Sequential( nn.Linear(hidden_dim, output_dim), nn.LayerNorm(output_dim), ) self.patch_proj = nn.Sequential( nn.Linear(token_dim, output_dim), nn.LayerNorm(output_dim), ) def _preprocess(self, x: torch.Tensor) -> torch.Tensor: if self.branch_type == "structure": gray = _grayscale(_blur(x, kernel_size=7)) return gray.repeat(1, 3, 1, 1) if self.branch_type == "texture": return torch.clamp((x - _blur(x, kernel_size=5)) * 0.5 + 0.5, 0.0, 1.0) if self.branch_type == "line": return _sobel_edges(x) if self.branch_type == "color": return _color_tone(x) raise ValueError(f"unknown branch_type: {self.branch_type}") def _prompt_pool(self, patch_tokens: torch.Tensor) -> torch.Tensor: if patch_tokens.size(1) == 0: return patch_tokens.new_zeros((patch_tokens.size(0), patch_tokens.size(-1))) prompt = F.normalize(self.token_prompt.float(), dim=0) patch_norm = F.normalize(patch_tokens.float(), dim=-1) attn_scores = torch.einsum("bnd,d->bn", patch_norm, prompt) attn = torch.softmax(attn_scores, dim=1).to(dtype=patch_tokens.dtype) return (patch_tokens * attn.unsqueeze(-1)).sum(dim=1) def _edge_pool(self, raw_view: torch.Tensor, patch_tokens: torch.Tensor, patch_hw: tuple[int, int]) -> torch.Tensor: if patch_tokens.size(1) == 0: return patch_tokens.new_zeros((patch_tokens.size(0), patch_tokens.size(-1))) edge_map = _grayscale(_sobel_edges(raw_view.float())) patch_weights = F.adaptive_avg_pool2d(edge_map, patch_hw).flatten(1) patch_weights = patch_weights / patch_weights.sum(dim=1, keepdim=True).clamp_min(1e-6) return (patch_tokens.float() * patch_weights.unsqueeze(-1)).sum(dim=1).to(dtype=patch_tokens.dtype) def _mean_pool(self, patch_tokens: torch.Tensor, fallback: torch.Tensor) -> torch.Tensor: if patch_tokens.size(1) == 0: return fallback return patch_tokens.mean(dim=1) def _storage_pool(self, storage_tokens: torch.Tensor, fallback: torch.Tensor) -> torch.Tensor: if storage_tokens.size(1) == 0: return fallback return storage_tokens.mean(dim=1) def _summarize_tokens( self, cls_tokens: torch.Tensor, patch_tokens: torch.Tensor, storage_tokens: torch.Tensor, raw_view: torch.Tensor, patch_hw: tuple[int, int], ) -> torch.Tensor: if self.branch_type == "structure": aux = self._mean_pool(patch_tokens, cls_tokens) elif self.branch_type == "texture": aux = self._prompt_pool(patch_tokens) elif self.branch_type == "line": aux = self._edge_pool(raw_view, patch_tokens, patch_hw) elif self.branch_type == "color": aux = self._storage_pool(storage_tokens, self._mean_pool(patch_tokens, cls_tokens)) else: raise ValueError(f"unknown branch_type: {self.branch_type}") return torch.cat([cls_tokens, aux], dim=-1) def forward(self, raw_views: List[torch.Tensor], view_features: List[Optional[Dict[str, torch.Tensor]]], view_mask: torch.Tensor) -> Dict[str, torch.Tensor]: encoded_views = [] batch_size = raw_views[0].size(0) for view_id, raw_view in enumerate(raw_views): view_present = view_mask[:, view_id].to(dtype=torch.bool) missing_features = self.missing_embeddings.weight[view_id].unsqueeze(0).expand(batch_size, -1) if view_present.any(): features = missing_features.clone() feature_pack = view_features[view_id] assert feature_pack is not None cls_tokens = feature_pack["cls"] patch_tokens = feature_pack["patch_tokens"] storage_tokens = feature_pack["storage_tokens"] patch_hw = feature_pack["patch_hw"] token_summary = self._summarize_tokens( cls_tokens=cls_tokens, patch_tokens=patch_tokens, storage_tokens=storage_tokens, raw_view=raw_view[view_present], patch_hw=patch_hw, ) token_features = self.token_proj(token_summary.float()) cue_input = self._preprocess(raw_view[view_present]) lowlevel_features = self.lowlevel_encoder(cue_input.float()) merged = self.merge(torch.cat([token_features, lowlevel_features], dim=-1)) merged = merged + self.view_embeddings.weight[view_id].unsqueeze(0) features[view_present] = merged.to(dtype=missing_features.dtype) else: features = missing_features encoded_views.append(features) stacked = torch.stack(encoded_views, dim=1) weights = _masked_softmax(self.view_gate(stacked).squeeze(-1), view_mask, dim=1) view_branch_embeddings = F.normalize(self.out_proj(stacked.float()), dim=-1) fused = (view_branch_embeddings * weights.unsqueeze(-1)).sum(dim=1) embedding = F.normalize(fused, dim=-1) return { "embedding": embedding, "view_weights": weights, "view_embeddings": stacked, "view_branch_embeddings": view_branch_embeddings, } class PrototypeClassifier(nn.Module): def __init__(self, num_classes: int, embed_dim: int, num_prototypes: int = 4, temperature: float = 0.1) -> None: super().__init__() self.num_classes = num_classes self.num_prototypes = num_prototypes self.temperature = temperature self.prototypes = nn.Parameter(torch.randn(num_classes, num_prototypes, embed_dim)) nn.init.normal_(self.prototypes, std=0.02) def forward(self, embedding: torch.Tensor) -> Dict[str, torch.Tensor]: normalized_embedding = F.normalize(embedding, dim=-1) normalized_prototypes = F.normalize(self.prototypes, dim=-1) similarities = torch.einsum("bd,cpd->bcp", normalized_embedding, normalized_prototypes) prototype_logits = torch.logsumexp(similarities / self.temperature, dim=-1) return { "embedding": normalized_embedding, "prototype_logits": prototype_logits, "prototype_similarities": similarities, "normalized_prototypes": normalized_prototypes, } class ArtistStyleModel(nn.Module): def __init__( self, num_classes: int, branch_hidden_dim: int = 384, branch_dim: int = 192, embedding_dim: int = 256, num_prototypes: int = 4, prototype_temperature: float = 0.1, arcface_scale: float = 30.0, view_dropout_prob: float = 0.15, branch_dropout_prob: float = 0.10, backbone_repo_dir: str = "third_party/dinov3", backbone_entrypoint: str = "dinov3_vits16", backbone_weights_path: str = "artifacts/pretrained/dinov3/dinov3_vits16_pretrain_lvd1689m-08c60483.pth", freeze_backbone: bool = True, backbone_unfreeze_last_n_blocks: int = 0, ) -> None: super().__init__() self.branch_names = ["structure", "texture", "line", "color"] if embedding_dim % len(self.branch_names) != 0: raise ValueError("embedding_dim must be divisible by the number of branches") self.view_dropout_prob = float(view_dropout_prob) self.branch_dropout_prob = float(branch_dropout_prob) self.per_branch_embedding_dim = embedding_dim // len(self.branch_names) self.backbone = DinoV3Backbone( repo_dir=backbone_repo_dir, entrypoint=backbone_entrypoint, weights_path=backbone_weights_path, freeze_backbone=freeze_backbone, unfreeze_last_n_blocks=backbone_unfreeze_last_n_blocks, ) self.branches = nn.ModuleDict( { name: SharedDinoViewBranch( branch_type=name, token_dim=self.backbone.hidden_size, hidden_dim=branch_hidden_dim, output_dim=branch_dim, ) for name in self.branch_names } ) self.branch_projectors = nn.ModuleDict( { name: nn.Sequential( nn.Linear(branch_dim, self.per_branch_embedding_dim), nn.LayerNorm(self.per_branch_embedding_dim), ) for name in self.branch_names } ) self.branch_gate = nn.Sequential( nn.Linear(branch_dim, branch_dim), nn.SiLU(inplace=True), nn.Linear(branch_dim, 1), ) self.embedding_norm = nn.Identity() self.arcface = ArcFaceHead( num_classes=num_classes, embed_dim=embedding_dim, scale=arcface_scale, ) self.classifier = PrototypeClassifier( num_classes=num_classes, embed_dim=embedding_dim, num_prototypes=num_prototypes, temperature=prototype_temperature, ) def _encode_views(self, views: List[torch.Tensor], view_mask: torch.Tensor) -> List[Optional[Dict[str, torch.Tensor]]]: features: List[Optional[Dict[str, torch.Tensor]]] = [] for view_id, view_tensor in enumerate(views): view_present = view_mask[:, view_id].to(dtype=torch.bool) if view_present.any(): features.append(self.backbone(view_tensor[view_present])) else: features.append(None) return features def encode(self, full: torch.Tensor, face: torch.Tensor, eye: torch.Tensor, view_mask: torch.Tensor | None = None) -> Dict[str, torch.Tensor]: views = [full, face, eye] if view_mask is None: view_mask = full.new_ones((full.size(0), 3)) effective_view_mask = _stochastic_keep_mask(view_mask, self.view_dropout_prob, self.training) view_features = self._encode_views(views, effective_view_mask) branch_embeddings = [] branch_view_weights = [] branch_view_embeddings = [] for branch_name in self.branch_names: branch_out = self.branches[branch_name](views, view_features=view_features, view_mask=effective_view_mask) branch_embeddings.append(branch_out["embedding"]) branch_view_weights.append(branch_out["view_weights"]) branch_view_embeddings.append(branch_out["view_branch_embeddings"]) branch_stack = torch.stack(branch_embeddings, dim=1) stacked_view_weights = torch.stack(branch_view_weights, dim=1) stacked_view_embeddings = torch.stack(branch_view_embeddings, dim=1) branch_mask = full.new_ones((full.size(0), len(self.branch_names))) effective_branch_mask = _stochastic_keep_mask(branch_mask, self.branch_dropout_prob, self.training) branch_weights = _masked_softmax(self.branch_gate(branch_stack).squeeze(-1), effective_branch_mask, dim=1) masked_branch_stack = branch_stack * effective_branch_mask.unsqueeze(-1) projected_chunks = [] for idx, branch_name in enumerate(self.branch_names): projected = F.normalize(self.branch_projectors[branch_name](branch_stack[:, idx].float()), dim=-1) projected_chunks.append(projected) branch_projected_stack = torch.stack(projected_chunks, dim=1) weighted_branch_projected = branch_projected_stack * branch_weights.unsqueeze(-1) * effective_branch_mask.unsqueeze(-1) raw_embedding = weighted_branch_projected.flatten(1) embedding = F.normalize(self.embedding_norm(raw_embedding.float()), dim=-1) return { "embedding": embedding, "branch_embeddings": branch_stack, "masked_branch_embeddings": masked_branch_stack, "branch_weights": branch_weights, "branch_projected_embeddings": branch_projected_stack, "weighted_branch_projected_embeddings": weighted_branch_projected, "view_weights": {name: stacked_view_weights[:, idx] for idx, name in enumerate(self.branch_names)}, "stacked_view_weights": stacked_view_weights, "stacked_view_embeddings": stacked_view_embeddings, "view_mask": view_mask, "effective_view_mask": effective_view_mask, "branch_mask": branch_mask, "effective_branch_mask": effective_branch_mask, } def forward(self, full: torch.Tensor, face: torch.Tensor, eye: torch.Tensor, view_mask: torch.Tensor | None = None) -> Dict[str, torch.Tensor]: encoded = self.encode(full, face, eye, view_mask=view_mask) arcface = self.arcface(encoded["embedding"]) classified = self.classifier(encoded["embedding"]) return { **encoded, **arcface, **classified, }