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, }