Files
modeling_zhaocui/final/model/mofe.py
T

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