from __future__ import annotations import torch from torch import nn class AlignedFusionModel(nn.Module): def __init__( self, kind: str, dims: tuple[int, int, int], steps: int = 50, hidden: int = 128, dropout: float = 0.15, ) -> None: super().__init__() if kind not in {"concat", "gate", "crossattn"}: raise ValueError(f"unknown model kind: {kind}") self.kind = kind self.hidden = hidden self.projections = nn.ModuleList( nn.Sequential(nn.Linear(size, hidden), nn.GELU(), nn.LayerNorm(hidden)) for size in dims ) self.position = nn.Parameter(torch.randn(1, steps, hidden) * 0.02) self.modality = nn.Parameter(torch.randn(1, 1, 3, hidden) * 0.02) self.dropout = nn.Dropout(dropout) if kind == "concat": self.fusion = nn.Sequential( nn.Linear(hidden * 3 + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout) ) elif kind == "gate": self.gate_score = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.Tanh(), nn.Linear(hidden // 2, 1)) self.fusion = nn.Sequential( nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout) ) else: layer = nn.TransformerEncoderLayer( d_model=hidden, nhead=4, dim_feedforward=hidden * 2, dropout=dropout, activation="gelu", batch_first=True, norm_first=True, ) self.cross_encoder = nn.TransformerEncoder(layer, num_layers=2, enable_nested_tensor=False) self.fusion = nn.Sequential( nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout) ) self.temporal = nn.GRU( input_size=hidden, hidden_size=hidden // 2, num_layers=1, batch_first=True, bidirectional=True, ) self.head = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout)) self.classifier = nn.Linear(hidden // 2, 3) self.regressor = nn.Linear(hidden // 2, 1) def forward(self, xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: torch.Tensor): masks = masks.bool() pos = self.position[:, :masks.shape[1]] encoded = [] for modality, (projection, x) in enumerate(zip(self.projections, xs)): token = projection(x) token = token + pos + self.modality[:, :, modality, :] token = token * masks[:, :, modality, None] encoded.append(token) stack = torch.stack(encoded, dim=2) # B x T x M x D availability = masks.to(stack.dtype) gate_weights = None if self.kind == "concat": fused = self.fusion(torch.cat((stack.flatten(2), availability), dim=-1)) elif self.kind == "gate": scores = self.gate_score(stack).squeeze(-1) scores = scores.masked_fill(~masks, -1e4) gate_weights = torch.softmax(scores, dim=-1) * availability gate_weights = gate_weights / gate_weights.sum(dim=-1, keepdim=True).clamp_min(1e-8) weighted = (stack * gate_weights[..., None]).sum(dim=2) fused = self.fusion(torch.cat((weighted, availability), dim=-1)) else: batch, steps, modalities, hidden = stack.shape flat = stack.reshape(batch, steps * modalities, hidden) valid = masks.reshape(batch, steps * modalities).clone() empty = ~valid.any(dim=1) if empty.any(): valid[empty, 0] = True flat[empty, 0] = 0.0 attended = self.cross_encoder(flat, src_key_padding_mask=~valid) attended = attended.reshape(batch, steps, modalities, hidden) observed_count = availability.sum(dim=2, keepdim=True) pooled = (attended * availability[..., None]).sum(dim=2) / observed_count.clamp_min(1.0) fused = self.fusion(torch.cat((pooled, availability), dim=-1)) temporal, _ = self.temporal(self.dropout(fused)) time_weight = masks.any(dim=-1).to(temporal.dtype) empty_time = time_weight.sum(dim=1, keepdim=True) <= 0 if empty_time.any(): time_weight[empty_time.squeeze(1), 0] = 1.0 pooled = (temporal * time_weight[..., None]).sum(dim=1) / time_weight.sum(dim=1, keepdim=True).clamp_min(1.0) hidden = self.head(pooled) logits = self.classifier(hidden) intensity = 3.0 * torch.tanh(self.regressor(hidden).squeeze(-1)) return {"logits": logits, "intensity": intensity, "gate": gate_weights}