557 lines
30 KiB
Python
557 lines
30 KiB
Python
"""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),
|
|
}
|