359 lines
14 KiB
Python
359 lines
14 KiB
Python
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
|