Train ATI-HO and finalize project outputs

This commit is contained in:
2026-09-26 16:05:44 +08:00
parent a86560da64
commit 9cdd604117
358 changed files with 10540 additions and 173 deletions
+358
View File
@@ -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
+35
View File
@@ -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
),
}