from __future__ import annotations import math from typing import Any import torch import torch.nn.functional as F from torch import nn from .ati_ho_config import ATIConfig PAIR_INDICES = ((0, 1), (0, 2), (1, 2)) PAIR_NAMES = ("TA", "TV", "AV") def _center_class_parameters(value: torch.Tensor) -> torch.Tensor: """Apply C to the three class logits while leaving magnitude parameters alone.""" logits = value[..., :3] logits = logits - logits.mean(dim=-1, keepdim=True) return torch.cat((logits, value[..., 3:]), dim=-1) class PrivateTemporalEncoder(nn.Module): """One modality-private projection, BiGRU(32 each way), and attention pool.""" def __init__(self, input_dim: int, hidden: int, gru_hidden: int) -> None: super().__init__() self.projection = nn.Sequential( nn.Linear(input_dim, hidden), nn.GELU(), nn.LayerNorm(hidden) ) self.temporal = nn.GRU( input_size=hidden, hidden_size=gru_hidden, num_layers=1, batch_first=True, bidirectional=True, ) self.pool_score = nn.Linear(hidden, 1) self.output_dim = 2 * gru_hidden def forward(self, x: torch.Tensor, observed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: observed = observed.bool() projected = self.projection(x) projected = projected * observed.unsqueeze(-1).to(projected.dtype) sequence, _ = self.temporal(projected) sequence = sequence * observed.unsqueeze(-1).to(sequence.dtype) scores = self.pool_score(torch.tanh(sequence)).squeeze(-1) scores = scores.masked_fill(~observed, torch.finfo(scores.dtype).min) has_any = observed.any(dim=1, keepdim=True) weights = torch.softmax(scores, dim=1) weights = torch.where(has_any, weights, torch.zeros_like(weights)) pooled = torch.sum(sequence * weights.unsqueeze(-1), dim=1) return sequence, pooled class MainEffectHead(nn.Module): def __init__(self, hidden: int) -> None: super().__init__() self.network = nn.Sequential(nn.Linear(hidden, hidden), nn.GELU(), nn.Linear(hidden, 5)) def forward(self, pooled: torch.Tensor) -> torch.Tensor: return _center_class_parameters(self.network(pooled)) class AnchoredPairBranch(nn.Module): """A pair reads only two private streams; its four-term anchor is explicit.""" def __init__(self, hidden: int, config: ATIConfig) -> None: super().__init__() self.low_rank_enabled = config.low_rank self.cross_attention_enabled = config.cross_attention self.rank = config.rank if self.low_rank_enabled: self.left_factor = nn.Linear(hidden, config.rank) self.right_factor = nn.Linear(hidden, config.rank) self.low_rank_out = nn.Linear(config.rank, 5, bias=False) else: self.left_factor = None self.right_factor = None self.low_rank_out = None if self.cross_attention_enabled: self.left_to_right = nn.MultiheadAttention( hidden, config.attention_heads, batch_first=True ) self.right_to_left = nn.MultiheadAttention( hidden, config.attention_heads, batch_first=True ) self.left_norm1 = nn.LayerNorm(hidden) self.right_norm1 = nn.LayerNorm(hidden) self.left_ffn = nn.Sequential( nn.Linear(hidden, config.attention_ffn), nn.GELU(), nn.Linear(config.attention_ffn, hidden), ) self.right_ffn = nn.Sequential( nn.Linear(hidden, config.attention_ffn), nn.GELU(), nn.Linear(config.attention_ffn, hidden), ) self.left_norm2 = nn.LayerNorm(hidden) self.right_norm2 = nn.LayerNorm(hidden) self.cross_out = nn.Linear(hidden * 2, 5, bias=False) else: self.left_to_right = None self.right_to_left = None self.left_norm1 = None self.right_norm1 = None self.left_ffn = None self.right_ffn = None self.left_norm2 = None self.right_norm2 = None self.cross_out = None init = min(max(config.eta_init, 1e-5), 1 - 1e-5) self.eta_logit = nn.Parameter(torch.tensor(math.log(init / (1.0 - init)))) # q(x0,y)=q(x,y0)=q(x0,y0)=offset. Four-term subtraction cancels it. # D0 deliberately leaves this offset in the output as a leakage control. self.anchor_offset = nn.Parameter(torch.zeros(5)) @staticmethod def _masked_mean(sequence: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: weights = mask.to(sequence.dtype).unsqueeze(-1) return (sequence * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0) @staticmethod def _safe_key_mask(mask: torch.Tensor) -> torch.Tensor: safe = mask.clone() empty = ~safe.any(dim=1) if empty.any(): safe[empty, 0] = True return safe def _core( self, left: torch.Tensor, right: torch.Tensor, left_mask: torch.Tensor, right_mask: torch.Tensor, ) -> torch.Tensor: joint = left_mask.bool() & right_mask.bool() values: list[torch.Tensor] = [] if self.low_rank_enabled: assert self.left_factor is not None and self.right_factor is not None assert self.low_rank_out is not None product = torch.tanh(self.left_factor(left)) * torch.tanh(self.right_factor(right)) values.append(self.low_rank_out(self._masked_mean(product, joint))) if self.cross_attention_enabled: assert self.left_to_right is not None and self.right_to_left is not None assert self.left_norm1 is not None and self.right_norm1 is not None assert self.left_ffn is not None and self.right_ffn is not None assert self.left_norm2 is not None and self.right_norm2 is not None assert self.cross_out is not None safe_left = self._safe_key_mask(left_mask.bool()) safe_right = self._safe_key_mask(right_mask.bool()) left_msg, _ = self.left_to_right( left, right, right, key_padding_mask=~safe_right, need_weights=False ) right_msg, _ = self.right_to_left( right, left, left, key_padding_mask=~safe_left, need_weights=False ) left_context = self.left_norm1(left + left_msg) right_context = self.right_norm1(right + right_msg) left_context = self.left_norm2(left_context + self.left_ffn(left_context)) right_context = self.right_norm2(right_context + self.right_ffn(right_context)) left_context = left_context * left_mask.unsqueeze(-1).to(left_context.dtype) right_context = right_context * right_mask.unsqueeze(-1).to(right_context.dtype) pooled = torch.cat( (self._masked_mean(left_context, joint), self._masked_mean(right_context, joint)), dim=-1, ) cross = self.cross_out(pooled) values.append(torch.sigmoid(self.eta_logit) * cross) if not values: return left.new_zeros((left.shape[0], 5)) # Each branch has a bias-free output and a joint-observation gate. Thus # core(x, y0)=core(x0, y)=core(x0, y0)=0 exactly. return torch.stack(values, dim=0).sum(dim=0) def forward( self, left: torch.Tensor, right: torch.Tensor, left_mask: torch.Tensor, right_mask: torch.Tensor, *, anchored: bool, ) -> torch.Tensor: raw_xy = self._core(left, right, left_mask, right_mask) + self.anchor_offset if anchored: # Four-term difference: # q(x,y)-q(x,x0)-q(x0,y)+q(x0,y0) = core(x,y). # The three absent-modality terms equal anchor_offset by the # joint gate and bias-free core, so they cancel algebraically. value = raw_xy - self.anchor_offset else: value = raw_xy return _center_class_parameters(value) class ATIHOModel(nn.Module): """Five-parameter additive multimodal predictor with exact modality anchors.""" def __init__(self, dims: tuple[int, int, int], config: ATIConfig, steps: int = 50) -> None: super().__init__() self.dims = tuple(int(d) for d in dims) self.steps = int(steps) self.config = config hidden = config.hidden self.encoders = nn.ModuleList( PrivateTemporalEncoder(dim, hidden, config.gru_hidden_per_direction) for dim in dims ) self.main_heads = nn.ModuleList(MainEffectHead(hidden) for _ in dims) self.mask_heads = nn.ModuleList(nn.Linear(hidden, 1) for _ in dims) self.pair_branches = nn.ModuleList( AnchoredPairBranch(hidden, config) for _ in PAIR_INDICES ) self.baseline = nn.Parameter(torch.zeros(5)) def forward( self, xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: torch.Tensor, *, return_details: bool = True, ) -> dict[str, Any]: if len(xs) != 3: raise ValueError("ATI–HO requires text, audio, and vision streams") if masks.ndim != 3 or masks.shape[-1] != 3: raise ValueError(f"masks must be B x T x 3, got {tuple(masks.shape)}") if masks.shape[1] > self.steps: raise ValueError(f"ATI–HO supports at most {self.steps} steps") masks = masks.bool() sequences: list[torch.Tensor] = [] pooled: list[torch.Tensor] = [] mask_logits: list[torch.Tensor] = [] main_effects: list[torch.Tensor] = [] for modality, (encoder, head, mask_head, x) in enumerate( zip(self.encoders, self.main_heads, self.mask_heads, xs) ): if x.shape[-1] != self.dims[modality]: raise ValueError( f"modality {modality} has {x.shape[-1]} features, expected {self.dims[modality]}" ) sequence, representation = encoder(x, masks[..., modality]) # Missing-mask baseline has a zero pooled representation. Explicit # subtraction makes every main effect zero at that baseline. baseline_raw = head(representation.new_zeros(representation.shape)) effect = _center_class_parameters(head(representation) - baseline_raw) sequences.append(sequence) pooled.append(representation) mask_logits.append(mask_head(sequence).squeeze(-1)) main_effects.append(effect) pair_effects: list[torch.Tensor] = [] pair_penalties: list[torch.Tensor] = [] for branch, (left_idx, right_idx) in zip(self.pair_branches, PAIR_INDICES): pair = branch( sequences[left_idx], sequences[right_idx], masks[..., left_idx], masks[..., right_idx], anchored=self.config.anchored, ) pair_effects.append(pair) pair_penalties.append(pair.square().mean()) main_tensor = torch.stack(main_effects, dim=1) pair_tensor = torch.stack(pair_effects, dim=1) params = self.baseline.unsqueeze(0) + main_tensor.sum(dim=1) + pair_tensor.sum(dim=1) logits = params[:, :3] probabilities = torch.softmax(logits, dim=-1) nu_negative = 3.0 * torch.sigmoid(params[:, 3]) nu_positive = 3.0 * torch.sigmoid(params[:, 4]) predicted_class = logits.argmax(dim=-1) hard_intensity = torch.where( predicted_class == 0, -nu_negative, torch.where(predicted_class == 2, nu_positive, torch.zeros_like(nu_positive)), ) soft_intensity = probabilities[:, 2] * nu_positive - probabilities[:, 0] * nu_negative result: dict[str, Any] = { "logits": logits, "probabilities": probabilities, "predicted_class": predicted_class, "intensity": hard_intensity, "soft_intensity": soft_intensity, "nu_negative": nu_negative, "nu_positive": nu_positive, "params": params, "interaction_penalty": torch.stack(pair_penalties).mean(), "mask_logits": torch.stack(mask_logits, dim=-1), } if return_details: result.update( { "baseline": self.baseline.unsqueeze(0).expand(xs[0].shape[0], -1), "main_effects": main_tensor, "pair_effects": pair_tensor, "main_sequences": torch.stack(sequences, dim=1), } ) return result def task_loss( output: dict[str, Any], y_cls: torch.Tensor, y_reg: torch.Tensor, *, lambda_interaction: float, lambda_mask: float, mask_target: torch.Tensor | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """CE + conditional polarity magnitude + low-weight continuous Huber.""" class_loss = F.cross_entropy(output["logits"], y_cls) negative = y_reg < 0 positive = y_reg > 0 target_mag = torch.abs(y_reg) / 3.0 magnitude_parts: list[torch.Tensor] = [] if negative.any(): magnitude_parts.append( F.smooth_l1_loss(output["nu_negative"][negative] / 3.0, target_mag[negative]) ) if positive.any(): magnitude_parts.append( F.smooth_l1_loss(output["nu_positive"][positive] / 3.0, target_mag[positive]) ) magnitude_loss = torch.stack(magnitude_parts).mean() if magnitude_parts else class_loss.new_zeros(()) continuous_loss = F.huber_loss( output["soft_intensity"] / 3.0, y_reg / 3.0, delta=0.25 ) interaction_loss = output["interaction_penalty"] mask_loss = class_loss.new_zeros(()) if lambda_mask > 0: if mask_target is None: raise ValueError("mask_target is required when the visibility-mask auxiliary loss is enabled") mask_loss = F.binary_cross_entropy_with_logits( output["mask_logits"], mask_target.to(output["mask_logits"].dtype) ) total = ( class_loss + magnitude_loss + 0.2 * continuous_loss + lambda_interaction * interaction_loss + lambda_mask * mask_loss ) parts = { "classification": class_loss, "conditional_magnitude": magnitude_loss, "continuous_huber": continuous_loss, "interaction": interaction_loss, "visibility_mask": mask_loss, "total": total, } return total, parts