from __future__ import annotations import torch from torch import nn class AlignedFusionModel(nn.Module): """Early concatenation + BiGRU model for the supplied aligned sequence.""" 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 != "concat": raise ValueError(f"only the selected EarlyConcat model is maintained; got: {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) self.fusion = nn.Sequential( nn.Linear(hidden * 3 + 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) fused = self.fusion(torch.cat((stack.flatten(2), 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) pooled = pooled / 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}