| 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)))) |
| |
| |
|
|
| 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") |
| |
| |
| 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() |
|
|