Files

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