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