Complete standalone final deliverable and unaligned Q2 results

This commit is contained in:
2026-09-25 22:22:37 +08:00
parent c6b018e5d0
commit adc9c2064b
267 changed files with 15479 additions and 7976 deletions
+23
View File
@@ -0,0 +1,23 @@
"""Q2 model entry points, one file per paper scheme."""
from .crg import CRG, StructuredGaussianImputer
from .c0 import C0
from .c1 import C1
from .c2 import C2
from .c3 import C3
from .c4 import C4
from .c5 import C5
from .c6 import C6
from .c6_no_distance import C6NoDistance
from .c6_no_reconstruction import C6NoReconstruction
from .c6_pointmask import C6PointMask
from .c7_distill import C7Distill
from .c7_group import C7Group
from .early_concat import AlignedFusionModel
from .mofe import MixtureOfFusionExperts
__all__ = [
"CRG", "StructuredGaussianImputer", "C0", "C1", "C2", "C3", "C4", "C5", "C6",
"C6NoDistance", "C6NoReconstruction", "C6PointMask", "C7Distill", "C7Group",
"AlignedFusionModel", "MixtureOfFusionExperts",
]
+57
View File
@@ -0,0 +1,57 @@
"""C0: observed statistics and mask baseline from E题V2, table 5.9."""
from __future__ import annotations
import numpy as np
from sklearn.linear_model import LogisticRegression, Ridge
MODALITIES = ("text", "audio", "vision")
def sample_statistics(arrays: dict[str, np.ndarray], mask: np.ndarray) -> np.ndarray:
"""Mean, standard deviation, missing fraction and longest gap per modality."""
mask = np.asarray(mask, dtype=bool)
if mask.ndim != 3 or mask.shape[-1] != 3:
raise ValueError("mask must have shape (N, T, 3)")
parts = []
for index, name in enumerate(MODALITIES):
x = np.asarray(arrays[name], dtype=np.float32)
if x.shape[:2] != mask.shape[:2]:
raise ValueError(f"{name}: feature and mask shapes disagree")
visible = mask[:, :, index]
count = visible.sum(axis=1, keepdims=True)
mean = (x * visible[:, :, None]).sum(axis=1) / np.maximum(count, 1)
variance = (((x - mean[:, None, :]) ** 2) * visible[:, :, None]).sum(axis=1) / np.maximum(count, 1)
missing = 1.0 - visible.mean(axis=1, keepdims=True)
max_gap = []
for row in visible:
longest = current = 0
for observed in row:
current = 0 if observed else current + 1
longest = max(longest, current)
max_gap.append(longest / max(len(row), 1))
parts.extend((mean, np.sqrt(variance), missing, np.asarray(max_gap, np.float32)[:, None]))
return np.concatenate(parts, axis=1).astype(np.float32)
class C0:
"""Logistic polarity classifier and Ridge intensity regressor."""
def __init__(self) -> None:
self.classifier = LogisticRegression(C=0.05, max_iter=2500, random_state=20260924)
self.regressor = Ridge(alpha=25.0)
def fit(self, arrays: dict[str, np.ndarray], mask: np.ndarray,
polarity: np.ndarray, intensity: np.ndarray) -> "C0":
features = sample_statistics(arrays, mask)
self.classifier.fit(features, polarity)
self.regressor.fit(features, intensity)
return self
def predict(self, arrays: dict[str, np.ndarray], mask: np.ndarray) -> dict[str, np.ndarray]:
features = sample_statistics(arrays, mask)
probabilities = np.zeros((len(features), 3), np.float64)
probabilities[:, self.classifier.classes_] = self.classifier.predict_proba(features)
return {
"probabilities": probabilities,
"intensity": np.clip(self.regressor.predict(features), -3.0, 3.0),
}
+9
View File
@@ -0,0 +1,9 @@
"""C1: masked BiGRU, without probabilistic completion or explicit gates."""
from .crg import CRG
class C1(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=False, use_joint_draws=False,
use_final_gate=False, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C2: Gaussian posterior mean completion, without trajectory integration."""
from .crg import CRG
class C2(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=False,
use_final_gate=False, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C3: joint trajectory integration and final reliability/content fusion."""
from .crg import CRG
class C3(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=False,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C4: C3 with bounded source attention and a null source."""
from .crg import CRG
class C4(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=False, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C5: C4 with reliability-modulated recurrent updates."""
from .crg import CRG
class C5(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=True, use_low_rank=False, **kwargs)
+9
View File
@@ -0,0 +1,9 @@
"""C6: C5 with the optional rank-four CP interaction residual enabled."""
from .crg import CRG
class C6(CRG):
def __init__(self, **kwargs):
super().__init__(use_imputer=True, use_joint_draws=True,
use_final_gate=True, use_source_attention=True,
reliability_update=True, use_low_rank=True, **kwargs)
+7
View File
@@ -0,0 +1,7 @@
"""C6 diagnostic: disable uncertainty distance and span penalties."""
from .c6 import C6
class C6NoDistance(C6):
def __init__(self, **kwargs):
super().__init__(reliability_hparams=(0.5, 0.05, 0.0, 0.0), **kwargs)
+8
View File
@@ -0,0 +1,8 @@
"""C6 diagnostic: omit auxiliary hidden-feature reconstruction while fitting."""
from .c6 import C6
class C6NoReconstruction(C6):
"""Use the C6 forward pass and set reconstruction loss weight to zero."""
reconstruction_loss_weight = 0.0
+8
View File
@@ -0,0 +1,8 @@
"""C6 diagnostic: train with independent point masks instead of spans."""
from .c6 import C6
class C6PointMask(C6):
"""Use the C6 forward pass with mask_kind='point' during fitting."""
training_mask_kind = "point"
+41
View File
@@ -0,0 +1,41 @@
"""C7 distillation-only branch: C6 architecture plus teacher loss."""
from __future__ import annotations
import math
import numpy as np
import torch
from torch.nn import functional as F
from .c6 import C6
DISTILL_TEMPERATURE = 2.0
DISTILL_WEIGHT = 0.1
class C7Distill(C6):
"""Inference uses C6; training adds weighted teacher distillation."""
def distillation_per_sample(student: dict, teacher: dict,
original: np.ndarray, current: np.ndarray) -> torch.Tensor:
"""Entropy/retention-weighted KL and score term for Q2 distillation."""
temp = DISTILL_TEMPERATURE
p_teacher = teacher["tempered_probs_by_path"].mean(dim=0).detach().clamp_min(1e-8)
p_student = student["tempered_probs_by_path"].mean(dim=0).clamp_min(1e-8)
entropy = -(p_teacher * p_teacher.log()).sum(dim=-1)
confidence_weight = (1.0 - entropy / math.log(3.0)).clamp(0.0, 1.0)
orig_t = torch.as_tensor(original, device=p_teacher.device, dtype=torch.float32)
curr_t = torch.as_tensor(current, device=p_teacher.device, dtype=torch.float32)
retained = []
for modality in range(3):
denominator = orig_t[:, :, modality].sum(dim=1)
ratio = (orig_t[:, :, modality] * curr_t[:, :, modality]).sum(dim=1) / denominator.clamp_min(1.0)
retained.append(torch.where(denominator > 0, ratio, torch.ones_like(ratio)))
weight = confidence_weight * torch.stack(retained, dim=-1).mean(dim=-1)
kl = (p_teacher * (p_teacher.log() - p_student.log())).sum(dim=-1) * temp * temp
teacher_score = teacher["mixed_score"].detach()
student_score = student["mixed_score"]
regression = F.huber_loss((teacher_score - student_score) / 3.0,
torch.zeros_like(teacher_score), reduction="none", delta=0.25)
return weight * (kl + regression)
+32
View File
@@ -0,0 +1,32 @@
"""C7 group-risk-only branch: C6 architecture plus smooth worst-group loss."""
from __future__ import annotations
import numpy as np
import torch
from .c6 import C6
class C7Group(C6):
"""Inference uses C6; training adds smooth worst-group risk."""
def smooth_group_risk(losses: torch.Tensor, group_ids: np.ndarray,
lambda_group: float = 0.1,
group_temperature: float = 0.05) -> torch.Tensor:
"""Match the selected group penalty from the Q2 training protocol."""
if group_temperature <= 0 or not 0 <= lambda_group <= 1:
raise ValueError("invalid group risk parameters")
groups = torch.as_tensor(group_ids, device=losses.device, dtype=torch.long)
if groups.shape != losses.shape:
raise ValueError("group_ids must match per-sample losses")
group_losses, priors = [], []
for group in torch.unique(groups):
selected = groups == group
group_losses.append(losses[selected].mean())
priors.append(selected.float().mean())
values = torch.stack(group_losses)
prior = torch.stack(priors).clamp_min(1e-8)
expected = (prior * values).sum()
worst = group_temperature * torch.logsumexp(torch.log(prior) + values / group_temperature, dim=0)
return (1.0 - lambda_group) * expected + lambda_group * worst
+556
View File
@@ -0,0 +1,556 @@
"""Structured Gaussian imputation and reliability-aware CRG sequence model."""
from __future__ import annotations
import math
from typing import Sequence
import torch
from torch import nn
from torch.nn import functional as F
MODALITIES = ("text", "audio", "vision")
INPUT_DIMS = (768, 74, 35)
HIDDEN = 32
SHARED_STATE = 8
PRIVATE_STATE = 4
STATE_DIM = SHARED_STATE + len(MODALITIES) * PRIVATE_STATE
def _inv_softplus(value: float) -> float:
return math.log(math.expm1(value))
class StructuredGaussianImputer(nn.Module):
"""Linear-Gaussian shared/private state model with exact block-Gaussian inference.
The state is [shared(8), text-private(4), audio-private(4), vision-private(4)].
Each modality emits from the shared state and its own private state only. The
filtering likelihood uses the matrix determinant lemma, retaining its log-det
normalization without forming a covariance matrix in observation space.
"""
def __init__(self, input_dims: Sequence[int] = INPUT_DIMS) -> None:
super().__init__()
self.input_dims = tuple(int(x) for x in input_dims)
self.state_dim = STATE_DIM
transition_mask = torch.zeros(STATE_DIM, STATE_DIM)
blocks = [slice(0, SHARED_STATE)] + [
slice(SHARED_STATE + i * PRIVATE_STATE, SHARED_STATE + (i + 1) * PRIVATE_STATE)
for i in range(len(MODALITIES))
]
for block in blocks:
transition_mask[block, block] = 1.0
self.register_buffer("transition_mask", transition_mask)
self.transition_raw = nn.Parameter(0.8 * torch.eye(STATE_DIM))
self.mu0 = nn.Parameter(torch.zeros(STATE_DIM))
self.pi0_raw = nn.Parameter(torch.full((STATE_DIM,), _inv_softplus(1.0)))
self.q_raw = nn.Parameter(torch.full((STATE_DIM,), _inv_softplus(0.08)))
self.emission_raw = nn.ParameterList()
self.biases = nn.ParameterList()
self.r_raw = nn.ParameterList()
for index, dim in enumerate(self.input_dims):
mask = torch.zeros(dim, STATE_DIM)
mask[:, :SHARED_STATE] = 1.0
private_start = SHARED_STATE + index * PRIVATE_STATE
mask[:, private_start:private_start + PRIVATE_STATE] = 1.0
self.register_buffer(f"emission_mask_{index}", mask)
self.emission_raw.append(nn.Parameter(torch.randn(dim, STATE_DIM) * 0.025))
self.biases.append(nn.Parameter(torch.zeros(dim)))
self.r_raw.append(nn.Parameter(torch.full((dim,), _inv_softplus(0.5))))
def _transition(self) -> torch.Tensor:
matrix = self.transition_raw * self.transition_mask
norm = torch.linalg.matrix_norm(matrix, ord=2).clamp_min(1e-8)
return matrix * torch.clamp(0.98 / norm, max=1.0)
def _covariances(self) -> tuple[torch.Tensor, torch.Tensor]:
eye = torch.eye(self.state_dim, device=self.mu0.device, dtype=self.mu0.dtype)
p0 = torch.diag(F.softplus(self.pi0_raw) + 1e-4) + 1e-5 * eye
q = torch.diag(F.softplus(self.q_raw) + 1e-4) + 1e-5 * eye
return p0, q
def emissions(self) -> list[torch.Tensor]:
return [raw * getattr(self, f"emission_mask_{i}") for i, raw in enumerate(self.emission_raw)]
def _filter(
self,
xs: Sequence[torch.Tensor],
observed: torch.Tensor,
*,
calculate_log_likelihood: bool,
retain_states: bool,
) -> tuple[torch.Tensor | None, dict[str, list[torch.Tensor]] | None]:
# xs[m]: [B,T,Dm], observed: [B,T,3]
batch, steps, _ = observed.shape
transition = self._transition()
p0, process_noise = self._covariances()
emissions = self.emissions()
noise = [F.softplus(x) + 1e-4 for x in self.r_raw]
mu_prior = self.mu0.expand(batch, -1)
p_prior = p0.expand(batch, -1, -1)
total_nll = torch.zeros(batch, device=observed.device, dtype=mu_prior.dtype)
prior_means: list[torch.Tensor] = []
prior_covs: list[torch.Tensor] = []
filtered_means: list[torch.Tensor] = []
filtered_covs: list[torch.Tensor] = []
for t in range(steps):
if retain_states:
prior_means.append(mu_prior)
prior_covs.append(p_prior)
p_chol = torch.linalg.cholesky(p_prior + 1e-6 * torch.eye(self.state_dim, device=p_prior.device))
p_inv = torch.cholesky_inverse(p_chol)
information_parts: list[torch.Tensor] = []
vector_parts: list[torch.Tensor] = []
quadratic_parts: list[torch.Tensor] = []
logdet_r = torch.zeros(batch, device=p_prior.device, dtype=p_prior.dtype)
n_observed = torch.zeros_like(logdet_r)
for m, (x, emission, variance) in enumerate(zip(xs, emissions, noise)):
active = observed[:, t, m].to(dtype=mu_prior.dtype)
weights = active[:, None] / variance[None, :]
centered = x[:, t] - self.biases[m]
residual = centered - mu_prior @ emission.T
information_parts.append(torch.einsum("di,bd,dj->bij", emission, weights, emission))
vector_parts.append((residual * weights) @ emission)
quadratic_parts.append((residual.square() * weights).sum(dim=-1))
logdet_r = logdet_r + active * torch.log(variance).sum()
n_observed = n_observed + active * x.shape[-1]
information = torch.stack(information_parts).sum(dim=0)
innovation = torch.stack(vector_parts).sum(dim=0)
precision = p_inv + information
precision_chol = torch.linalg.cholesky(precision + 1e-6 * torch.eye(self.state_dim, device=precision.device))
p_filtered = torch.cholesky_inverse(precision_chol)
mu_filtered = mu_prior + torch.einsum("bij,bj->bi", p_filtered, innovation)
if calculate_log_likelihood:
logdet_p = 2.0 * torch.log(torch.diagonal(p_chol, dim1=-2, dim2=-1)).sum(dim=-1)
logdet_precision = 2.0 * torch.log(torch.diagonal(precision_chol, dim1=-2, dim2=-1)).sum(dim=-1)
quad = torch.stack(quadratic_parts).sum(dim=0)
correction = torch.einsum("bi,bij,bj->b", innovation, p_filtered, innovation)
log_likelihood = logdet_r + logdet_p + logdet_precision + (quad - correction).clamp_min(0.0)
log_likelihood = log_likelihood + n_observed * math.log(2.0 * math.pi)
total_nll = total_nll + 0.5 * log_likelihood
if retain_states:
filtered_means.append(mu_filtered)
filtered_covs.append(p_filtered)
mu_prior = mu_filtered @ transition.T
p_prior = transition @ p_filtered @ transition.T + process_noise
states = None
if retain_states:
states = {
"prior_mean": prior_means,
"prior_cov": prior_covs,
"filtered_mean": filtered_means,
"filtered_cov": filtered_covs,
"transition": [transition],
}
return (total_nll if calculate_log_likelihood else None), states
def observed_nll(self, xs: Sequence[torch.Tensor], observed: torch.Tensor) -> torch.Tensor:
"""Exact observed-data Gaussian NLL, including covariance log determinants."""
nll, _ = self._filter(xs, observed, calculate_log_likelihood=True, retain_states=False)
assert nll is not None
return nll
@staticmethod
def _draw(mean: torch.Tensor, covariance: torch.Tensor, paths: int) -> torch.Tensor:
chol = torch.linalg.cholesky(covariance + 1e-5 * torch.eye(covariance.shape[-1], device=covariance.device))
noise = torch.randn((paths, *mean.shape), dtype=mean.dtype, device=mean.device)
return mean.unsqueeze(0) + torch.einsum("bij,kbj->kbi", chol, noise)
@torch.no_grad()
def complete(
self,
xs: Sequence[torch.Tensor],
observed: torch.Tensor,
paths: int,
*,
joint_draws: bool,
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
"""RTS smooth, draw joint latent trajectories, then draw missing emissions."""
_, stored = self._filter(xs, observed, calculate_log_likelihood=False, retain_states=True)
assert stored is not None
fm, fc = stored["filtered_mean"], stored["filtered_cov"]
pm, pc = stored["prior_mean"], stored["prior_cov"]
transition = stored["transition"][0]
steps = len(fm)
smoother_gains: list[torch.Tensor] = [torch.empty(0, device=observed.device)] * max(0, steps - 1)
smooth_cov: list[torch.Tensor] = [torch.empty(0, device=observed.device)] * steps
smooth_cov[-1] = fc[-1]
for t in range(steps - 2, -1, -1):
next_chol = torch.linalg.cholesky(pc[t + 1] + 1e-6 * torch.eye(self.state_dim, device=observed.device))
gain = torch.cholesky_solve((fc[t] @ transition.T).transpose(-1, -2), next_chol).transpose(-1, -2)
smoother_gains[t] = gain
smooth_cov[t] = fc[t] + gain @ (smooth_cov[t + 1] - pc[t + 1]) @ gain.transpose(-1, -2)
smooth_cov[t] = 0.5 * (smooth_cov[t] + smooth_cov[t].transpose(-1, -2))
if joint_draws:
state = torch.empty((paths, observed.shape[0], steps, self.state_dim), device=observed.device, dtype=fm[0].dtype)
state[:, :, -1] = self._draw(fm[-1], fc[-1], paths)
for t in range(steps - 2, -1, -1):
gain = smoother_gains[t]
conditional_mean = fm[t].unsqueeze(0) + torch.einsum(
"bij,kbj->kbi", gain, state[:, :, t + 1] - pm[t + 1].unsqueeze(0)
)
conditional_cov = fc[t] - gain @ pc[t + 1] @ gain.transpose(-1, -2)
conditional_cov = 0.5 * (conditional_cov + conditional_cov.transpose(-1, -2))
chol = torch.linalg.cholesky(conditional_cov + 1e-5 * torch.eye(self.state_dim, device=observed.device))
eps = torch.randn_like(conditional_mean)
state[:, :, t] = conditional_mean + torch.einsum("bij,kbj->kbi", chol, eps)
else:
means = torch.stack(fm, dim=1)
covs = torch.stack(smooth_cov, dim=1)
smoothed_means = [fm[-1]] * steps
smoothed_means[-1] = fm[-1]
for t in range(steps - 2, -1, -1):
smoothed_means[t] = fm[t] + torch.einsum(
"bij,bj->bi", smoother_gains[t], smoothed_means[t + 1] - pm[t + 1]
)
state = torch.stack(smoothed_means, dim=1).unsqueeze(0).expand(paths, -1, -1, -1)
completed: list[torch.Tensor] = []
variances: list[torch.Tensor] = []
for m, (x, emission) in enumerate(zip(xs, self.emissions())):
mean = torch.einsum("kbti,di->kbtd", state, emission) + self.biases[m]
if joint_draws:
noise = torch.randn_like(mean) * torch.sqrt(F.softplus(self.r_raw[m]) + 1e-4)
draws = mean + noise
else:
draws = mean
visible = observed[:, :, m].unsqueeze(0).unsqueeze(-1)
completed.append(torch.where(visible, x.unsqueeze(0), draws))
projected_cov = torch.einsum("di,btij,dj->btd", emission, torch.stack(smooth_cov, dim=1), emission)
variance = projected_cov + (F.softplus(self.r_raw[m]) + 1e-4)
variances.append(torch.where(observed[:, :, m, None], torch.zeros_like(variance), variance.clamp_min(1e-6)))
return completed, variances
class ReliabilityGRU(nn.Module):
"""One-layer BiGRU with directional time decay and rho-scaled updates."""
def __init__(self, input_dim: int, hidden: int = 16) -> None:
super().__init__()
self.hidden = hidden
self.x_proj = nn.Linear(input_dim, 3 * hidden)
self.h_proj = nn.Linear(hidden, 2 * hidden, bias=False)
self.candidate_h = nn.Linear(hidden, hidden, bias=False)
self.decay_raw = nn.Parameter(torch.full((hidden,), -3.0))
def _one_direction(
self,
x: torch.Tensor,
rho: torch.Tensor,
distance: torch.Tensor,
reverse: bool,
reliability_update: bool,
) -> torch.Tensor:
batch, steps, _ = x.shape
state = torch.zeros(batch, self.hidden, dtype=x.dtype, device=x.device)
x_parts = self.x_proj(x).chunk(3, dim=-1)
output: list[torch.Tensor | None] = [None] * steps
indices = range(steps - 1, -1, -1) if reverse else range(steps)
for t in indices:
if reliability_update:
decay = torch.exp(-F.softplus(self.decay_raw)[None, :] * distance[:, t:t + 1])
decayed_state = decay * state
else:
decayed_state = state
hz, hr = self.h_proj(decayed_state).chunk(2, dim=-1)
z = torch.sigmoid(x_parts[0][:, t] + hz)
r = torch.sigmoid(x_parts[1][:, t] + hr)
candidate = torch.tanh(x_parts[2][:, t] + self.candidate_h(r * decayed_state))
effective_z = rho[:, t:t + 1] * z if reliability_update else z
state = (1.0 - effective_z) * decayed_state + effective_z * candidate
output[t] = state
return torch.stack([v for v in output if v is not None], dim=1)
def forward(
self,
x: torch.Tensor,
rho: torch.Tensor,
dminus: torch.Tensor,
dplus: torch.Tensor,
reliability_update: bool,
) -> torch.Tensor:
if not reliability_update:
rho = torch.ones_like(rho)
return torch.cat((
self._one_direction(x, rho, dminus, False, reliability_update),
self._one_direction(x, rho, dplus, True, reliability_update),
), dim=-1)
class CRG(nn.Module):
"""Quality-aware multimodal sequence predictor for a configured ablation."""
def __init__(
self,
imputer: StructuredGaussianImputer | None = None,
input_dims: Sequence[int] = INPUT_DIMS,
*,
use_imputer: bool = True,
use_joint_draws: bool = True,
use_final_gate: bool = True,
use_source_attention: bool = True,
reliability_update: bool = True,
use_low_rank: bool = True,
reliability_hparams: tuple[float, float, float, float] = (0.5, 0.05, 0.05, 0.05),
) -> None:
super().__init__()
self.use_imputer = use_imputer
self.use_joint_draws = use_joint_draws
self.use_final_gate = use_final_gate
self.use_source_attention = use_source_attention
self.reliability_update = reliability_update
self.use_low_rank = use_low_rank
self.imputer = imputer if imputer is not None else StructuredGaussianImputer(input_dims)
self.projections = nn.ModuleList(
nn.Sequential(nn.Linear(d, HIDDEN), nn.LayerNorm(HIDDEN), nn.GELU()) for d in input_dims
)
recurrent_input = HIDDEN + 13
self.temporal = nn.ModuleList(ReliabilityGRU(recurrent_input, 16) for _ in MODALITIES)
rho_imp, lambda_u, lambda_gap, lambda_span = reliability_hparams
if not 0.0 < rho_imp < 1.0 or min(lambda_u, lambda_gap, lambda_span) < 0.0:
raise ValueError("reliability requires 0<rho_imp<1 and nonnegative distance/uncertainty penalties")
self.register_buffer("rho_imp", torch.tensor(float(rho_imp)))
self.register_buffer("rel_u", torch.full((3,), float(lambda_u)))
self.register_buffer("rel_gap", torch.full((3,), float(lambda_gap)))
self.register_buffer("rel_span", torch.full((3,), float(lambda_span)))
self.query = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.key = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.value = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.relative_bias = nn.Embedding(99, 1)
nn.init.zeros_(self.relative_bias.weight)
self.cross_base = nn.Linear(HIDDEN, HIDDEN)
self.cross_out = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.cross_eta_logit = nn.Parameter(torch.tensor(-1.0))
self.content_score = nn.Sequential(nn.Linear(HIDDEN, 16), nn.Tanh(), nn.Linear(16, 1, bias=False))
self.null_expert = nn.Parameter(torch.zeros(HIDDEN))
self.pool_hidden = nn.Linear(HIDDEN, 16)
self.pool_score = nn.Linear(16, 1, bias=False)
self.reconstruction_heads = nn.ModuleList(nn.Linear(HIDDEN, d) for d in input_dims)
self.cp_factors = nn.ModuleList(nn.Linear(HIDDEN + 1, 4, bias=False) for _ in MODALITIES)
self.cp_output = nn.Parameter(torch.randn(4, HIDDEN) * 0.02)
self.low_rank_output = nn.Linear(HIDDEN, HIDDEN, bias=False)
self.low_rank_eta_logit = nn.Parameter(torch.tensor(-4.0))
self.head = nn.Sequential(nn.Linear(HIDDEN + 18, 64), nn.GELU(), nn.Dropout(0.2))
self.classifier = nn.Linear(64, 3)
self.magnitude_mean = nn.Linear(64, 2)
self.concentration_raw = nn.Parameter(torch.full((2,), _inv_softplus(6.0)))
@staticmethod
def _gap_features(observed: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
# Time positions are valid sequence locations even when all three sources are missing.
batch, steps, modalities = observed.shape
device = observed.device
positions = torch.arange(steps, device=device).view(1, steps).expand(batch, -1)
previous = torch.full((batch, modalities), -1, device=device, dtype=torch.long)
before, before_edge = [], []
for t in range(steps):
before_edge.append(previous < 0)
before.append(torch.where(previous < 0, torch.ones_like(previous, dtype=torch.float32), (t - previous).float() / max(1, steps - 1)))
previous = torch.where(observed[:, t], torch.full_like(previous, t), previous)
following = torch.full((batch, modalities), steps, device=device, dtype=torch.long)
after, after_edge = [None] * steps, [None] * steps
for t in range(steps - 1, -1, -1):
after_edge[t] = following >= steps
after[t] = torch.where(following >= steps, torch.ones_like(following, dtype=torch.float32), (following - t).float() / max(1, steps - 1))
following = torch.where(observed[:, t], torch.full_like(following, t), following)
dminus = torch.stack(before, dim=1)
dplus = torch.stack([x for x in after if x is not None], dim=1)
edge_minus = torch.stack(before_edge, dim=1)
edge_plus = torch.stack([x for x in after_edge if x is not None], dim=1)
dminus = torch.where(observed, torch.zeros_like(dminus), dminus)
dplus = torch.where(observed, torch.zeros_like(dplus), dplus)
edge_minus = edge_minus & ~observed
edge_plus = edge_plus & ~observed
missing = ~observed
left_run = torch.zeros((batch, steps, modalities), device=device, dtype=torch.float32)
run = torch.zeros((batch, modalities), device=device, dtype=torch.float32)
for t in range(steps):
run = torch.where(missing[:, t], run + 1.0, torch.zeros_like(run))
left_run[:, t] = run
right_run = torch.zeros_like(left_run)
run.zero_()
for t in range(steps - 1, -1, -1):
run = torch.where(missing[:, t], run + 1.0, torch.zeros_like(run))
right_run[:, t] = run
span = torch.where(missing, (left_run + right_run - 1.0) / max(1, steps), torch.zeros_like(left_run))
return dminus, dplus, span, torch.stack((edge_minus, edge_plus), dim=-1).float()
def _reliability(
self, observed: torch.Tensor, uncertainty: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
dminus, dplus, span, edges = self._gap_features(observed)
gap = torch.minimum(dminus, dplus)
gap = torch.where(observed, torch.zeros_like(gap), gap)
u = torch.where(observed, torch.zeros_like(uncertainty), uncertainty).clamp_min(0.0)
qstar = observed.float() # External Q2 quality is unavailable: q*=1 only for visible rows; J=0.
rho_missing = self.rho_imp.clamp(1e-4, 0.999) * torch.exp(
-self.rel_u[None, None, :] * u
-self.rel_gap[None, None, :] * gap
-self.rel_span[None, None, :] * span
)
rho = torch.where(observed, qstar, rho_missing).clamp(1e-4, 1.0)
return rho, u, gap, span, dminus, dplus, torch.cat((qstar.unsqueeze(-1), torch.zeros_like(qstar).unsqueeze(-1), edges), dim=-1)
def _cross_source(self, hidden: torch.Tensor, rho: torch.Tensor) -> torch.Tensor:
# hidden [B,T,M,H]; each query reads every legal time in each other source.
batch, steps, modalities, width = hidden.shape
outputs = []
q = self.query(hidden)
k = self.key(hidden)
v = torch.tanh(self.value(hidden))
loc = torch.arange(steps, device=hidden.device)
relative_index = (loc[None, :] - loc[:, None] + 49).clamp(0, 98)
relative = self.relative_bias(relative_index).squeeze(-1)
for target in range(modalities):
numerator = torch.zeros((batch, steps, width), device=hidden.device, dtype=hidden.dtype)
denominator = torch.ones((batch, steps, 1), device=hidden.device, dtype=hidden.dtype)
for source in range(modalities):
if source == target:
continue
raw = torch.matmul(q[:, :, target], k[:, :, source].transpose(-1, -2)) / math.sqrt(width)
scores = 2.0 * torch.tanh(raw + relative)
base = 1.0 / (max(1, modalities - 1) * steps)
weights = base * rho[:, None, :, source] * torch.exp(scores.clamp(-2.0, 2.0))
numerator = numerator + torch.matmul(weights, v[:, :, source])
denominator = denominator + weights.sum(dim=-1, keepdim=True)
context = numerator / denominator
eta = torch.sigmoid(self.cross_eta_logit)
outputs.append(torch.tanh(self.cross_base(hidden[:, :, target]) + eta * self.cross_out(context)))
return torch.stack(outputs, dim=2)
def _low_rank_residual(self, gated: torch.Tensor) -> torch.Tensor:
# Linear CP factors use [1; z_m] and subtract their constant all-zero term.
batch, steps, modalities, width = gated.shape
one = torch.ones((batch, steps, 1), device=gated.device, dtype=gated.dtype)
products = torch.ones((batch, steps, 4), device=gated.device, dtype=gated.dtype)
constant = torch.ones(4, device=gated.device, dtype=gated.dtype)
for m in range(modalities):
factor = self.cp_factors[m](torch.cat((one, gated[:, :, m]), dim=-1))
products = products * factor
zero_input = torch.zeros((1, 1, width + 1), device=gated.device, dtype=gated.dtype)
zero_input[..., 0] = 1.0
constant = constant * self.cp_factors[m](zero_input)[0, 0]
residual = (products - constant) @ self.cp_output
return torch.sigmoid(self.low_rank_eta_logit) * self.low_rank_output(torch.tanh(residual))
def forward(
self,
xs: Sequence[torch.Tensor],
observed_mask: torch.Tensor,
*,
paths: int = 4,
joint_draws: bool | None = None,
) -> dict[str, torch.Tensor | list[torch.Tensor]]:
batch, steps, modalities = observed_mask.shape
if joint_draws is None:
joint_draws = self.use_joint_draws
if self.use_imputer:
completed, variance = self.imputer.complete(xs, observed_mask, paths, joint_draws=joint_draws)
else:
completed = [torch.where(observed_mask[:, :, m, None], x, torch.zeros_like(x)).unsqueeze(0) for m, x in enumerate(xs)]
variance = [torch.zeros_like(x) for x in xs]
paths = completed[0].shape[0]
uncertainty_parts = [v.mean(dim=-1) for v in variance]
uncertainty = torch.stack(uncertainty_parts, dim=-1)
rho, u, gap, span, dminus, dplus, quality_fields = self._reliability(observed_mask, uncertainty)
position = torch.linspace(0.0, 1.0, steps, device=observed_mask.device, dtype=xs[0].dtype)
pe = torch.stack((torch.sin(2 * math.pi * position), torch.cos(2 * math.pi * position),
torch.sin(4 * math.pi * position), torch.cos(4 * math.pi * position)), dim=-1)
encoded_paths: list[torch.Tensor] = []
reconstructed_paths: list[list[torch.Tensor]] = []
logits_paths: list[torch.Tensor] = []
beta_paths: list[torch.Tensor] = []
fusion_weight_paths: list[torch.Tensor] = []
null_weight_paths: list[torch.Tensor] = []
time_pool_weight_paths: list[torch.Tensor] = []
for path_index in range(paths):
enc = [projection(completed[m][path_index]) for m, projection in enumerate(self.projections)]
hmods, reconstruction = [], []
for m, encoder in enumerate(self.temporal):
if self.use_final_gate or self.use_source_attention or self.reliability_update:
scalar = torch.cat((observed_mask[:, :, m:m + 1].float(), quality_fields[:, :, m],
torch.log1p(u[:, :, m:m + 1]), dminus[:, :, m:m + 1],
dplus[:, :, m:m + 1], span[:, :, m:m + 1],
pe.unsqueeze(0).expand(batch, -1, -1)), dim=-1)
else:
# C1/C2 receive only the visibility mask and legal position code.
scalar = torch.zeros((batch, steps, 13), dtype=pe.dtype, device=pe.device)
scalar[:, :, 0] = observed_mask[:, :, m].float()
scalar[:, :, -4:] = pe.unsqueeze(0)
# q*, J_Q, edge flags, directional gaps, uncertainty and span are explicit.
seq = torch.cat((enc[m], scalar), dim=-1)
h = encoder(seq, rho[:, :, m], dminus[:, :, m], dplus[:, :, m], self.reliability_update)
hmods.append(h)
reconstruction.append(self.reconstruction_heads[m](h))
hidden = torch.stack(hmods, dim=2)
if self.use_source_attention:
enhanced = self._cross_source(hidden, rho)
else:
enhanced = hidden
if self.use_final_gate:
content = 2.0 * torch.tanh(self.content_score(enhanced).squeeze(-1))
weights_unnorm = rho * torch.exp(content.clamp(-2.0, 2.0))
denom = 1.0 + weights_unnorm.sum(dim=-1, keepdim=True)
alpha = weights_unnorm / denom
null_alpha = 1.0 / denom.squeeze(-1)
gated = enhanced * alpha.unsqueeze(-1)
fused = gated.sum(dim=2) + null_alpha.unsqueeze(-1) * self.null_expert
else:
if self.use_imputer:
alpha = torch.full_like(observed_mask.float(), 1.0 / modalities)
else:
alpha = observed_mask.float() / observed_mask.float().sum(dim=-1, keepdim=True).clamp_min(1.0)
null_alpha = torch.zeros((batch, steps), device=observed_mask.device, dtype=alpha.dtype)
fused = (enhanced * alpha.unsqueeze(-1)).sum(dim=2)
gated = enhanced * alpha.unsqueeze(-1)
if self.use_low_rank:
fused = fused + self._low_rank_residual(gated)
pool_logits = 2.0 * torch.tanh(self.pool_score(torch.tanh(self.pool_hidden(fused))).squeeze(-1))
pool_weight = torch.softmax(pool_logits, dim=1)
pooled = (pool_weight.unsqueeze(-1) * fused).sum(dim=1)
missing_rate = 1.0 - observed_mask.float().mean(dim=1)
mean_rho = rho.mean(dim=1) if (self.use_final_gate or self.use_source_attention or self.reliability_update) else observed_mask.float().mean(dim=1)
max_gap = gap.max(dim=1).values
max_span = span.max(dim=1).values
edge_rate = self._gap_features(observed_mask)[3].mean(dim=1).reshape(batch, -1)
stats = torch.cat((missing_rate, mean_rho, max_gap, max_span, edge_rate), dim=-1)
representation = torch.cat((pooled, stats), dim=-1)
feature = self.head(representation)
logits_paths.append(self.classifier(feature))
mean_fraction = torch.sigmoid(self.magnitude_mean(feature)).clamp(1e-4, 1.0 - 1e-4)
concentration = F.softplus(self.concentration_raw).clamp_min(1e-3)
alpha_beta = torch.stack((mean_fraction * concentration, (1.0 - mean_fraction) * concentration), dim=-1)
beta_paths.append(alpha_beta)
reconstructed_paths.append(reconstruction)
encoded_paths.append(hidden)
fusion_weight_paths.append(alpha)
null_weight_paths.append(null_alpha)
time_pool_weight_paths.append(pool_weight)
class_logits = torch.stack(logits_paths, dim=0)
beta_params = torch.stack(beta_paths, dim=0)
class_probs_by_path = torch.softmax(class_logits, dim=-1)
beta_mean = beta_params[..., 0] / beta_params.sum(dim=-1)
conditional_mean = 3.0 * (class_probs_by_path[..., 2] * beta_mean[..., 1] - class_probs_by_path[..., 0] * beta_mean[..., 0])
return {
"class_logits": class_logits,
"class_probs_by_path": class_probs_by_path,
"class_probs": class_probs_by_path.mean(dim=0),
"tempered_probs_by_path": torch.softmax(class_logits / 2.0, dim=-1),
"beta_params": beta_params,
"beta_mean": beta_mean,
"mixed_score": conditional_mean.mean(dim=0),
"reconstructions": [torch.stack([reconstructed_paths[k][m] for k in range(paths)], dim=0) for m in range(modalities)],
"reliability": rho,
"imputation_uncertainty": uncertainty,
"gap": gap,
"span": span,
"distance_before": dminus,
"distance_after": dplus,
"fusion_weights_by_path": torch.stack(fusion_weight_paths, dim=0),
"null_weights_by_path": torch.stack(null_weight_paths, dim=0),
"time_pool_weights_by_path": torch.stack(time_pool_weight_paths, dim=0),
"low_rank_scale": torch.sigmoid(self.low_rank_eta_logit),
}
+67
View File
@@ -0,0 +1,67 @@
from __future__ import annotations
import torch
from torch import nn
class AlignedFusionModel(nn.Module):
"""Early concatenation + BiGRU model for the supplied aligned sequence."""
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 != "concat":
raise ValueError(f"only the selected EarlyConcat model is maintained; got: {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)
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,
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)
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)
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}
+224
View File
@@ -0,0 +1,224 @@
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,
}