Files
modeling_zhaocui/submit/final/model/ati_ho.py
T

359 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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