Complete standalone final deliverable and unaligned Q2 results
This commit is contained in:
@@ -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",
|
||||
]
|
||||
@@ -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),
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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),
|
||||
}
|
||||
@@ -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}
|
||||
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user