from __future__ import annotations import math import torch from torch import Tensor, nn from torch.nn import functional as F def masked_mean_prototype(per_image: Tensor, mask: Tensor) -> Tensor: if per_image.ndim != 3 or mask.shape != per_image.shape[:2]: raise ValueError("expected per_image [batch,references,width] and matching mask") mask = mask.to(device=per_image.device, dtype=torch.bool) counts = mask.sum(dim=1) if torch.any(counts < 1) or torch.any(counts > 8): raise ValueError("each reference set must contain 1-8 images") weights = mask.unsqueeze(-1).to(per_image.dtype) mean = (per_image * weights).sum(dim=1) / counts.unsqueeze(-1) return F.normalize(mean, dim=-1) class StyleMixerBlock(nn.Module): def __init__(self, width: int, heads: int, dropout: float = 0.0) -> None: super().__init__() self.query_norm = nn.LayerNorm(width) self.token_norm = nn.LayerNorm(width) self.attention = nn.MultiheadAttention(width, heads, batch_first=True) self.attention_dropout = nn.Dropout(dropout) self.output_norm = nn.LayerNorm(width) self.mlp = nn.Sequential( nn.Linear(width, 4 * width), nn.GELU(), nn.Dropout(dropout), nn.Linear(4 * width, width), nn.Dropout(dropout), ) def forward(self, query: Tensor, tokens: Tensor, token_mask: Tensor) -> Tensor: value, _ = self.attention( self.query_norm(query), self.token_norm(tokens), self.token_norm(tokens), key_padding_mask=~token_mask, need_weights=False, ) query = query + self.attention_dropout(value) return query + self.mlp(self.output_norm(query)) class UnifiedStyleEncoder(nn.Module): """Encode cached full/face vision tokens and optional Anima descriptors.""" def __init__( self, *, external_dim: int = 1152, external_tokens: int = 30, anima_dim: int = 4096, anima_tokens: int = 3, width: int = 384, blocks: int = 2, heads: int = 6, embedding_dim: int = 512, use_anima: bool = True, dropout: float = 0.0, reference_grounding: bool = False, functional_factors: bool = False, ) -> None: super().__init__() if ( blocks < 1 or width % heads or (reference_grounding and embedding_dim % heads) or not 0.0 <= dropout < 1.0 ): raise ValueError("blocks must be positive and width must be divisible by heads") if functional_factors and embedding_dim != 512: raise ValueError("functional factorization currently requires a 512-D embedding") self.external_tokens = external_tokens self.anima_tokens = anima_tokens if use_anima else 0 self.external_projector = nn.Sequential( nn.LayerNorm(external_dim), nn.Linear(external_dim, width), ) self.anima_projector = ( nn.Sequential(nn.LayerNorm(anima_dim), nn.Linear(anima_dim, width)) if use_anima else None ) self.full_positions = nn.Parameter(torch.randn(external_tokens, width) * 0.02) self.face_positions = nn.Parameter(torch.randn(external_tokens, width) * 0.02) self.anima_positions = ( nn.Parameter(torch.randn(anima_tokens, width) * 0.02) if use_anima else None ) self.reference_grounding = reference_grounding self.functional_factors = functional_factors query_count = 5 if functional_factors else 1 self.style_query = nn.Parameter(torch.randn(1, query_count, width) * 0.02) self.blocks = nn.ModuleList( [StyleMixerBlock(width, heads, dropout) for _ in range(blocks)] ) if functional_factors: self.factor_output = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, 96)) self.shared_output = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, 128)) self.output = nn.Sequential(nn.LayerNorm(512), nn.Linear(512, embedding_dim)) self.factor_decoders = nn.ModuleList([nn.Linear(96, external_dim) for _ in range(4)]) else: self.output = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, embedding_dim)) if reference_grounding: self.consensus_query = nn.Parameter(torch.randn(1, 1, embedding_dim) * 0.02) self.consensus_attention = nn.MultiheadAttention( embedding_dim, heads, batch_first=True, dropout=dropout ) self.consensus_condition = nn.Linear(embedding_dim, width) self.consensus_gamma = nn.Parameter(torch.tensor(-2.9444)) def _read_tokens( self, tokens: Tensor, token_mask: Tensor, conditioning: Tensor | None = None ) -> tuple[Tensor, Tensor | None]: query = self.style_query.expand(tokens.shape[0], -1, -1) if conditioning is not None: query = query + self.consensus_condition(conditioning)[:, None] for block in self.blocks: query = block(query, tokens, token_mask) if self.functional_factors: factors = F.normalize(self.factor_output(query[:, :4]), dim=-1) shared = self.shared_output(query[:, 4]) embedding = self.output(torch.cat((factors.flatten(1), shared), dim=-1)) return F.normalize(embedding, dim=-1), factors return F.normalize(self.output(query[:, 0]), dim=-1), None def encode_images( self, full_features: Tensor, *, face_features: Tensor | None = None, face_mask: Tensor | None = None, anima_features: Tensor | None = None, conditioning: Tensor | None = None, ) -> Tensor: """Encode independent images without padding them into reference sets.""" if full_features.ndim != 3 or full_features.shape[1] != self.external_tokens: raise ValueError("full_features must have shape [images,external_tokens,width]") images = full_features.shape[0] full = self.external_projector(full_features) + self.full_positions tokens = [full] masks = [torch.ones(full.shape[:2], dtype=torch.bool, device=full.device)] if face_features is not None: if face_features.shape != full_features.shape: raise ValueError("face_features must match full_features") if face_mask is None or face_mask.shape != (images,): raise ValueError("face_mask is required and must contain one value per image") face = self.external_projector(face_features) tokens.append(face + self.face_positions) masks.append(face_mask[:, None].to(torch.bool).expand(-1, self.external_tokens)) if self.anima_projector is not None: expected = (images, self.anima_tokens) if anima_features is None or anima_features.shape[:2] != expected: raise ValueError("anima_features are required by this encoder configuration") anima = self.anima_projector(anima_features) tokens.append(anima + self.anima_positions) masks.append(torch.ones(anima.shape[:2], dtype=torch.bool, device=anima.device)) token_tensor = torch.cat(tokens, dim=1) token_mask = torch.cat(masks, dim=1) return self._read_tokens(token_tensor, token_mask, conditioning)[0] def encode_images_with_factors( self, full_features: Tensor, *, face_features: Tensor | None = None, face_mask: Tensor | None = None, anima_features: Tensor | None = None, conditioning: Tensor | None = None, ) -> tuple[Tensor, Tensor | None]: if full_features.ndim != 3 or full_features.shape[1] != self.external_tokens: raise ValueError("full_features must have shape [images,external_tokens,width]") images = full_features.shape[0] full = self.external_projector(full_features) + self.full_positions tokens = [full] masks = [torch.ones(full.shape[:2], dtype=torch.bool, device=full.device)] if face_features is not None: if face_features.shape != full_features.shape: raise ValueError("face_features must match full_features") if face_mask is None or face_mask.shape != (images,): raise ValueError("face_mask is required and must contain one value per image") tokens.append(self.external_projector(face_features) + self.face_positions) masks.append(face_mask[:, None].to(torch.bool).expand(-1, self.external_tokens)) if self.anima_projector is not None: if anima_features is None or anima_features.shape[:2] != (images, self.anima_tokens): raise ValueError("anima_features are required by this encoder configuration") tokens.append(self.anima_projector(anima_features) + self.anima_positions) masks.append(torch.ones(anima_features.shape[:2], dtype=torch.bool, device=full.device)) return self._read_tokens(torch.cat(tokens, dim=1), torch.cat(masks, dim=1), conditioning) def encode_reference_sets( self, full_features: Tensor, reference_mask: Tensor, *, face_features: Tensor | None = None, face_mask: Tensor | None = None, anima_features: Tensor | None = None, ) -> tuple[Tensor, Tensor, Tensor | None]: batch, references = full_features.shape[:2] flat_face = None if face_features is None else face_features.flatten(0, 1) flat_face_mask = None if face_mask is None else face_mask.flatten() flat_anima = None if anima_features is None else anima_features.flatten(0, 1) per_image, factors = self.encode_images_with_factors( full_features.flatten(0, 1), face_features=flat_face, face_mask=flat_face_mask, anima_features=flat_anima, ) per_image = per_image.reshape(batch, references, -1) factors = None if factors is None else factors.reshape(batch, references, 4, -1) base = masked_mean_prototype(per_image, reference_mask) if not self.reference_grounding: return base, per_image, factors consensus, _ = self.consensus_attention( self.consensus_query.expand(batch, -1, -1), per_image, per_image, key_padding_mask=~reference_mask, need_weights=False, ) consensus = F.normalize(consensus[:, 0], dim=-1) reread, factors = self.encode_images_with_factors( full_features.flatten(0, 1), face_features=flat_face, face_mask=flat_face_mask, anima_features=flat_anima, conditioning=consensus[:, None].expand(-1, references, -1).reshape(batch * references, -1), ) reread = reread.reshape(batch, references, -1) factors = None if factors is None else factors.reshape(batch, references, 4, -1) agreement = F.cosine_similarity(reread, consensus[:, None], dim=-1) centered = agreement - ( (agreement * reference_mask).sum(1, keepdim=True) / reference_mask.sum(1, keepdim=True).clamp_min(1) ) weights = (1.0 + 0.25 * torch.tanh(4.0 * centered)) * reference_mask grounded = F.normalize((reread * weights[..., None]).sum(1) / weights.sum(1, keepdim=True), dim=-1) gamma = torch.sigmoid(self.consensus_gamma) return F.normalize(base + gamma * (grounded - base), dim=-1), reread, factors def functional_response_loss(self, factor_codes: Tensor, full_features: Tensor) -> Tensor: """Reconstruct artist response after removing a shared prompt/seed cell mean.""" if not self.functional_factors: return full_features.new_zeros(()) groups = ((2, 3, 4, 5), (0, 6, 24), tuple(range(8, 24)), (1, 7, 25)) losses = [] for index, token_indices in enumerate(groups): teacher = full_features[..., list(token_indices), :].float().mean(dim=-2) teacher = F.normalize(teacher - teacher.mean(dim=1, keepdim=True), dim=-1) prediction = F.normalize(self.factor_decoders[index](factor_codes[..., index, :]).float(), dim=-1) losses.append((1.0 - F.cosine_similarity(prediction, teacher, dim=-1)).mean()) return torch.stack(losses).mean() def forward( self, full_features: Tensor, reference_mask: Tensor, *, face_features: Tensor | None = None, face_mask: Tensor | None = None, anima_features: Tensor | None = None, ) -> tuple[Tensor, Tensor]: if full_features.ndim != 4 or full_features.shape[2] != self.external_tokens: raise ValueError("full_features must have shape [batch,references,external_tokens,width]") batch, references = full_features.shape[:2] if reference_mask.shape != (batch, references): raise ValueError("reference_mask shape does not match full_features") prototype, per_image, _ = self.encode_reference_sets( full_features, reference_mask, face_features=face_features, face_mask=face_mask, anima_features=anima_features, ) return prototype, per_image class AngularPrototypeLoss(nn.Module): """Episode prototype softmax with an optional ArcFace-style positive margin.""" def __init__( self, initial_scale: float = 10.0, initial_bias: float | None = None ) -> None: super().__init__() if initial_scale <= 0: raise ValueError("initial_scale must be positive") self.raw_scale = nn.Parameter(torch.tensor(math.log(math.expm1(initial_scale)))) # initial_bias remains accepted for feasibility-script compatibility. A # common softmax bias is a no-op and is not a trainable production parameter. def forward( self, queries: Tensor, prototypes: Tensor, targets: Tensor, *, margin: float = 0.0, ) -> tuple[Tensor, Tensor]: if queries.ndim < 2 or prototypes.ndim != queries.ndim: raise ValueError("queries and prototypes must have matching [episode,...,width] ranks") if queries.shape[:-2] != prototypes.shape[:-2] or queries.shape[-1] != prototypes.shape[-1]: raise ValueError("queries and prototypes must have matching episode axes and widths") if targets.shape != queries.shape[:-1]: raise ValueError("targets must contain one class index per query") # Autocast can round a near-perfect BF16 cosine to exactly one. ArcFace's # sqrt(1-cos²) derivative is then singular, so margin math and CE stay FP32. cosine = ( F.normalize(queries, dim=-1) @ F.normalize(prototypes, dim=-1).transpose(-1, -2) ).float() loss_cosine = cosine if margin: if not 0.0 <= margin < math.pi / 2: raise ValueError("margin must be in [0, pi/2)") positive = cosine.gather(-1, targets.unsqueeze(-1)).clamp(-1 + 1e-6, 1 - 1e-6) positive_margin = positive * math.cos(margin) - torch.sqrt(1 - positive.square()) * math.sin(margin) loss_cosine = cosine.scatter(-1, targets.unsqueeze(-1), positive_margin) scale = F.softplus(self.raw_scale).clamp_max(100.0) loss_logits = scale * loss_cosine metric_logits = scale * cosine return ( F.cross_entropy(loss_logits.flatten(0, -2), targets.flatten()), metric_logits, ) def subset_consistency_loss(first: Tensor, second: Tensor) -> Tensor: if first.shape != second.shape or first.ndim != 2: raise ValueError("subset prototypes must have equal [batch,width] shapes") return (1.0 - F.cosine_similarity(first, second, dim=-1)).mean()