Files
modeling_zhaocui/deep_learning/Q2/q2/models.py
T

68 lines
2.7 KiB
Python

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}