Train ATI-HO and finalize project outputs
This commit is contained in:
@@ -0,0 +1,358 @@
|
||||
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
|
||||
@@ -0,0 +1,35 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict, dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ATIConfig:
|
||||
name: str
|
||||
low_rank: bool = False
|
||||
cross_attention: bool = False
|
||||
anchored: bool = True
|
||||
lambda_interaction: float = 1e-3
|
||||
lambda_mask: float = 0.0
|
||||
rank: int = 4
|
||||
hidden: int = 64
|
||||
gru_hidden_per_direction: int = 32
|
||||
attention_heads: int = 4
|
||||
attention_ffn: int = 128
|
||||
eta_init: float = 0.1
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
CONFIGS: dict[str, ATIConfig] = {
|
||||
"A0": ATIConfig(name="A0_main_effects"),
|
||||
"A1": ATIConfig(name="A1_low_rank_pairs", low_rank=True),
|
||||
"A2": ATIConfig(name="A2_anchored_pairwise", low_rank=True, cross_attention=True),
|
||||
"A3": ATIConfig(
|
||||
name="A3_pairwise_mask_aux", low_rank=True, cross_attention=True, lambda_mask=0.05
|
||||
),
|
||||
"D0": ATIConfig(
|
||||
name="D0_unanchored_diagnostic", low_rank=True, cross_attention=True, anchored=False
|
||||
),
|
||||
}
|
||||
Reference in New Issue
Block a user