整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,109 @@
|
||||
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}
|
||||
Reference in New Issue
Block a user