提交其余项目实验变更

This commit is contained in:
2026-09-25 10:41:58 +08:00
parent 83ec3d1a83
commit 95bd34599b
119 changed files with 5877 additions and 1709 deletions
+11 -53
View File
@@ -5,6 +5,8 @@ from torch import nn
class AlignedFusionModel(nn.Module):
"""Early concatenation + BiGRU model for the supplied aligned sequence."""
def __init__(
self,
kind: str,
@@ -14,8 +16,8 @@ class AlignedFusionModel(nn.Module):
dropout: float = 0.15,
) -> None:
super().__init__()
if kind not in {"concat", "gate", "crossattn"}:
raise ValueError(f"unknown model kind: {kind}")
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(
@@ -25,31 +27,9 @@ class AlignedFusionModel(nn.Module):
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.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,
@@ -72,38 +52,16 @@ class AlignedFusionModel(nn.Module):
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))
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) / time_weight.sum(dim=1, keepdim=True).clamp_min(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, "gate": gate_weights}
return {"logits": logits, "intensity": intensity}