225 lines
8.8 KiB
Python
225 lines
8.8 KiB
Python
from __future__ import annotations
|
|
|
|
from typing import Any
|
|
|
|
import torch
|
|
import torch.nn.functional as F
|
|
from torch import nn
|
|
|
|
|
|
SUBSETS: dict[str, tuple[int, ...]] = {
|
|
"T": (0,),
|
|
"A": (1,),
|
|
"V": (2,),
|
|
"TA": (0, 1),
|
|
"TV": (0, 2),
|
|
"AV": (1, 2),
|
|
"TAV": (0, 1, 2),
|
|
}
|
|
EXPERT_NAMES = tuple(SUBSETS)
|
|
EXPERT_BITS = {
|
|
name: tuple(int(i in indices) for i in range(3))
|
|
for name, indices in SUBSETS.items()
|
|
}
|
|
|
|
|
|
class MixtureOfFusionExperts(nn.Module):
|
|
"""Seven-subset, hard-availability MoFE with the selected MLP router.
|
|
|
|
Each modality has a private projection. Experts only receive the private
|
|
projections belonging to their subset. The weighted result is passed
|
|
through one shared temporal backbone and one shared prediction head.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
dims: tuple[int, int, int],
|
|
router: str = "mlp",
|
|
expert_names: tuple[str, ...] = EXPERT_NAMES,
|
|
availability_mode: str = "hard",
|
|
steps: int = 50,
|
|
latent_dim: int = 64,
|
|
hidden: int = 128,
|
|
dropout: float = 0.15,
|
|
) -> None:
|
|
super().__init__()
|
|
if router != "mlp":
|
|
raise ValueError(f"only the selected MLP router is maintained; got: {router}")
|
|
if availability_mode != "hard":
|
|
raise ValueError(f"only hard availability masking is maintained; got: {availability_mode}")
|
|
if tuple(expert_names) != EXPERT_NAMES:
|
|
raise ValueError("the selected MoFE uses all seven modality-subset experts")
|
|
|
|
self.dims = dims
|
|
self.router_kind = router
|
|
self.expert_names = tuple(expert_names)
|
|
self.availability_mode = availability_mode
|
|
self.steps = steps
|
|
self.latent_dim = latent_dim
|
|
self.hidden = hidden
|
|
|
|
# These projections are private to each modality and are not tied.
|
|
self.private_projections = nn.ModuleList(
|
|
nn.Sequential(nn.Linear(size, latent_dim), nn.GELU()) for size in dims
|
|
)
|
|
self.experts = nn.ModuleDict()
|
|
for name in self.expert_names:
|
|
n_modalities = len(SUBSETS[name])
|
|
self.experts[name] = nn.Sequential(
|
|
nn.Linear(n_modalities * latent_dim, hidden),
|
|
nn.GELU(),
|
|
nn.Dropout(dropout),
|
|
nn.Linear(hidden, latent_dim),
|
|
nn.LayerNorm(latent_dim),
|
|
)
|
|
|
|
router_input_dim = 9
|
|
self.router = nn.Sequential(
|
|
nn.Linear(router_input_dim, 16),
|
|
nn.GELU(),
|
|
nn.Linear(16, len(self.expert_names)),
|
|
)
|
|
|
|
# Shared early-fusion projection, BiGRU, and task heads.
|
|
self.all_missing_token = nn.Parameter(torch.zeros(1, 1, latent_dim))
|
|
self.input_projection = nn.Sequential(
|
|
nn.Linear(latent_dim + 3, hidden),
|
|
nn.GELU(),
|
|
nn.LayerNorm(hidden),
|
|
nn.Dropout(dropout),
|
|
)
|
|
self.dropout = 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)
|
|
|
|
@staticmethod
|
|
def _availability(masks: torch.Tensor, names: tuple[str, ...]) -> torch.Tensor:
|
|
masks = masks.bool()
|
|
columns = [masks[..., list(SUBSETS[name])].all(dim=-1) for name in names]
|
|
return torch.stack(columns, dim=-1)
|
|
|
|
def _router_features(
|
|
self,
|
|
private: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
|
masks: torch.Tensor,
|
|
) -> torch.Tensor:
|
|
observed = masks.to(dtype=private[0].dtype)
|
|
magnitude = torch.stack(
|
|
[torch.sqrt(x.square().mean(dim=-1) + 1e-8) for x in private], dim=-1
|
|
)
|
|
local_ratio = F.avg_pool1d(
|
|
observed.transpose(1, 2), kernel_size=5, stride=1, padding=2, count_include_pad=False
|
|
).transpose(1, 2)
|
|
return torch.cat((observed, torch.log1p(magnitude), local_ratio), dim=-1)
|
|
|
|
def _route(
|
|
self,
|
|
router_features: torch.Tensor,
|
|
availability: torch.Tensor,
|
|
force_expert: str | None,
|
|
) -> torch.Tensor:
|
|
scores = self.router(router_features)
|
|
scores = scores.masked_fill(~availability, -1e4)
|
|
weights = torch.softmax(scores, dim=-1) * availability.to(scores.dtype)
|
|
# In the full seven-expert model this is exactly the all-modalities-
|
|
# missing case. It also safely handles ablations with no eligible set.
|
|
has_expert = availability.any(dim=-1, keepdim=True)
|
|
weights = weights * has_expert.to(weights.dtype)
|
|
weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
|
|
|
|
if force_expert is not None:
|
|
if force_expert not in self.expert_names:
|
|
raise ValueError(f"expert {force_expert} is not enabled in this model")
|
|
expert_idx = self.expert_names.index(force_expert)
|
|
forced = torch.zeros_like(weights)
|
|
forced[..., expert_idx] = 1.0
|
|
# Force the requested expert where its modality subset is present;
|
|
# where it is unavailable, use the learned router over eligible
|
|
# experts instead of replacing observed information with zeros.
|
|
return torch.where(availability[..., expert_idx, None], forced, weights)
|
|
|
|
return weights
|
|
|
|
def forward(
|
|
self,
|
|
xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
|
|
masks: torch.Tensor,
|
|
force_expert: str | None = None,
|
|
) -> dict[str, Any]:
|
|
masks = masks.bool()
|
|
if masks.ndim != 3 or masks.shape[-1] != 3:
|
|
raise ValueError(f"masks must have shape B x T x 3, got {tuple(masks.shape)}")
|
|
if masks.shape[1] > self.steps:
|
|
raise ValueError(f"sequence has {masks.shape[1]} steps, model supports {self.steps}")
|
|
|
|
private_values = []
|
|
for modality, (projector, x) in enumerate(zip(self.private_projections, xs)):
|
|
projected = projector(x)
|
|
projected = projected * masks[..., modality, None].to(projected.dtype)
|
|
private_values.append(projected)
|
|
private = tuple(private_values)
|
|
router_features = self._router_features(private, masks)
|
|
availability = self._availability(masks, self.expert_names)
|
|
|
|
local_expert_outputs = []
|
|
for name in self.expert_names:
|
|
indices = SUBSETS[name]
|
|
expert_input = torch.cat([private[i] for i in indices], dim=-1)
|
|
local_expert_outputs.append(self.experts[name](expert_input))
|
|
expert_stack = torch.stack(local_expert_outputs, dim=-2)
|
|
|
|
alpha_local = self._route(router_features, availability, force_expert)
|
|
fused = (expert_stack * alpha_local[..., None]).sum(dim=-2)
|
|
has_expert = availability.any(dim=-1)
|
|
fused = torch.where(
|
|
has_expert[..., None], fused, self.all_missing_token.expand_as(fused)
|
|
)
|
|
|
|
# Restore a stable seven-column interface for saved diagnostics,
|
|
# including expert-set ablations.
|
|
alpha = masks.new_zeros((*masks.shape[:2], len(EXPERT_NAMES)), dtype=private[0].dtype)
|
|
expert_outputs = private[0].new_zeros((*masks.shape[:2], len(EXPERT_NAMES), self.latent_dim))
|
|
for local_idx, name in enumerate(self.expert_names):
|
|
global_idx = EXPERT_NAMES.index(name)
|
|
alpha[..., global_idx] = alpha_local[..., local_idx]
|
|
expert_outputs[..., global_idx, :] = expert_stack[..., local_idx, :]
|
|
|
|
fused_with_masks = torch.cat((fused, masks.to(fused.dtype)), dim=-1)
|
|
encoded = self.input_projection(fused_with_masks)
|
|
temporal, _ = self.temporal(self.dropout(encoded))
|
|
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))
|
|
|
|
bits = torch.tensor(
|
|
[EXPERT_BITS[name] for name in EXPERT_NAMES],
|
|
dtype=alpha.dtype,
|
|
device=alpha.device,
|
|
)
|
|
utility = torch.einsum("bte,em->btm", alpha, bits)
|
|
return {
|
|
"logits": logits,
|
|
"intensity": intensity,
|
|
"fused": fused,
|
|
"alpha": alpha,
|
|
"utility": utility,
|
|
"availability": availability,
|
|
"expert_outputs": expert_outputs,
|
|
"fallback": ~has_expert,
|
|
"router_features": router_features,
|
|
}
|