from __future__ import annotations from collections.abc import Mapping, Sequence import numpy as np import torch import torch.nn.functional as F from sklearn.dummy import DummyClassifier from sklearn.linear_model import LogisticRegression, Ridge from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error from sklearn.preprocessing import StandardScaler from torch import Tensor, nn from .alignment import make_block_mask from .losses import cross_modal_contrastive_loss from .metrics import retrieval_metrics from .types import MODALITIES class RetrievalProjection(nn.Module): """Equal-capacity linear heads for the cross-modal retrieval probe.""" def __init__(self, dimensions: Mapping[str, int], output_size: int = 128) -> None: super().__init__() self.projections = nn.ModuleDict( {name: nn.Linear(dimensions[name], output_size, bias=False) for name in MODALITIES} ) def forward(self, aligned: Mapping[str, Tensor]) -> dict[str, Tensor]: return { name: F.normalize(self.projections[name](aligned[name]), dim=-1) for name in MODALITIES } class MaskedCrossModalDecoder(nn.Module): """Reconstruct one missing modality from the other aligned streams.""" def __init__(self, dimensions: Mapping[str, int], hidden_size: int = 256) -> None: super().__init__() self.dimensions = dict(dimensions) input_size = sum(dimensions.values()) + len(MODALITIES) self.decoders = nn.ModuleDict( { target: nn.Sequential( nn.Linear(input_size, hidden_size), nn.GELU(), nn.Dropout(0.1), nn.Linear(hidden_size, dimensions[target]), ) for target in MODALITIES } ) def forward(self, aligned: Mapping[str, Tensor], target: str, mask: Tensor) -> Tensor: values = [] availability = [] for name in MODALITIES: present = torch.ones_like(mask, dtype=aligned[name].dtype) source = aligned[name] if name == target: present = (~mask).to(source.dtype) source = source.masked_fill(mask.unsqueeze(-1), 0.0) values.append(source) availability.append(present.unsqueeze(-1)) inputs = torch.cat((*values, *availability), dim=-1) return self.decoders[target](inputs) def _stack_aligned( sample_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], device: torch.device ) -> dict[str, Tensor]: return { name: torch.as_tensor( np.stack([aligned_by_id[sample_id][name] for sample_id in sample_ids]), dtype=torch.float32, device=device, ) for name in MODALITIES } def run_retrieval_probe( train_ids: Sequence[str], val_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], *, device: torch.device, seed: int, epochs: int = 20, batch_size: int = 16, learning_rate: float = 1e-3, temperature: float = 0.1, ) -> list[dict[str, float | str]]: """Fit modality projections on training clips, then score held-out retrieval. The positive is a matching common-grid index within a clip. This measures representation consistency; it is not independent temporal ground truth. """ if not train_ids or not val_ids: raise ValueError("retrieval probe needs non-empty train and validation sets") device_gen = torch.Generator(device=device) device_gen.manual_seed(seed) torch.manual_seed(seed) train = _stack_aligned(train_ids, aligned_by_id, device) validation = _stack_aligned(val_ids, aligned_by_id, device) dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES} model = RetrievalProjection(dimensions).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) rng = np.random.default_rng(seed) model.train() for _ in range(epochs): order = rng.permutation(len(train_ids)) for start in range(0, len(order), batch_size): indices = torch.as_tensor(order[start : start + batch_size], device=device) batch = {name: value.index_select(0, indices) for name, value in train.items()} projected = model(batch) loss = cross_modal_contrastive_loss(projected, temperature=temperature) optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() model.eval() with torch.no_grad(): projected = model(validation) directions = (("text", "audio"), ("audio", "text"), ("text", "vision"), ("vision", "text"), ("audio", "vision"), ("vision", "audio")) rows: list[dict[str, float | str]] = [] for query_name, target_name in directions: values = retrieval_metrics(projected[query_name], projected[target_name]) rows.append({"direction": f"{query_name}_to_{target_name}", **values}) return rows def run_within_clip_temporal_retrieval_probe( train_ids: Sequence[str], val_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], *, device: torch.device, seed: int, epochs: int = 20, batch_size: int = 16, learning_rate: float = 1e-3, temperature: float = 0.1, tolerance: int = 1, top_k: int = 3, ) -> list[dict[str, float | str]]: """Fit train-only cross-modal projections, then retrieve slots within each clip. Unlike global grid retrieval, each query's candidates are restricted to the target modality from that same held-out clip. A result is correct if its slot is within ``tolerance`` of the query slot. """ if not train_ids or not val_ids: raise ValueError("temporal retrieval needs non-empty train and validation sets") if tolerance < 0 or top_k < 1: raise ValueError("tolerance must be non-negative and top_k positive") torch.manual_seed(seed) train = _stack_aligned(train_ids, aligned_by_id, device) validation = _stack_aligned(val_ids, aligned_by_id, device) dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES} model = RetrievalProjection(dimensions).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) rng = np.random.default_rng(seed) model.train() for _ in range(epochs): order = rng.permutation(len(train_ids)) for start in range(0, len(order), batch_size): indices = torch.as_tensor(order[start : start + batch_size], device=device) batch = {name: value.index_select(0, indices) for name, value in train.items()} projected = model(batch) loss = cross_modal_contrastive_loss(projected, temperature=temperature) optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() model.eval() with torch.no_grad(): projected = model(validation) directions = ( ("text", "audio"), ("audio", "text"), ("text", "vision"), ("vision", "text"), ("audio", "vision"), ("vision", "audio"), ) rows: list[dict[str, float | str]] = [] for query_name, target_name in directions: distances_top1 = [] hits_top1 = [] hits_topk = [] for clip_index in range(len(val_ids)): query = projected[query_name][clip_index] target = projected[target_name][clip_index] scores = query @ target.T count = scores.shape[0] k = min(top_k, count) candidates = scores.topk(k=k, dim=-1).indices slots = torch.arange(count, device=device)[:, None] distances = (candidates - slots).abs() distances_top1.append(distances[:, 0].float()) hits_top1.append((distances[:, 0] <= tolerance).float()) hits_topk.append((distances <= tolerance).any(dim=-1).float()) top1_distance = torch.cat(distances_top1) rows.append( { "direction": f"{query_name}_to_{target_name}", "r_at_1": float(torch.cat(hits_top1).mean().item()), "r_at_3": float(torch.cat(hits_topk).mean().item()), "mase_slots": float(top1_distance.mean().item()), "exact_r_at_1": float((top1_distance == 0).float().mean().item()), "tolerance_slots": float(tolerance), "candidate_slots_per_clip": float(projected[query_name].shape[1]), "queries": float(len(val_ids) * projected[query_name].shape[1]), } ) return rows def _fixed_block_mask( count: int, grid_size: int, ratio: float, device: torch.device, salt: int ) -> Tensor: block = min(max(1, round(grid_size * ratio)), grid_size - 1) starts = torch.tensor( [(index * 17 + salt * 13) % (grid_size - block + 1) for index in range(count)], dtype=torch.long, device=device, ) offsets = torch.arange(block, device=device) mask = torch.zeros(count, grid_size, dtype=torch.bool, device=device) mask[torch.arange(count, device=device)[:, None], starts[:, None] + offsets] = True return mask def run_reconstruction_probe( train_ids: Sequence[str], val_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], *, device: torch.device, seed: int, ratio: float = 0.2, epochs: int = 25, batch_size: int = 16, learning_rate: float = 1e-3, ) -> list[dict[str, float | str]]: """Train the same decoder family on frozen alignments and score held-out clips.""" if not train_ids or not val_ids: raise ValueError("reconstruction probe needs non-empty train and validation sets") torch.manual_seed(seed) generator = torch.Generator(device=device) generator.manual_seed(seed) train = _stack_aligned(train_ids, aligned_by_id, device) validation = _stack_aligned(val_ids, aligned_by_id, device) dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES} model = MaskedCrossModalDecoder(dimensions).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) rng = np.random.default_rng(seed) model.train() for _ in range(epochs): order = rng.permutation(len(train_ids)) for start in range(0, len(order), batch_size): indices = torch.as_tensor(order[start : start + batch_size], device=device) batch = {name: value.index_select(0, indices) for name, value in train.items()} target_losses = [] for target in MODALITIES: mask = make_block_mask( len(indices), batch[target].shape[1], ratio, device, generator=generator, ) prediction = model(batch, target, mask) target_losses.append(F.smooth_l1_loss(prediction[mask], batch[target][mask])) loss = torch.stack(target_losses).mean() optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() model.eval() rows: list[dict[str, float | str]] = [] with torch.no_grad(): for target_index, target in enumerate(MODALITIES): mask = _fixed_block_mask( len(val_ids), validation[target].shape[1], ratio, device, target_index ) prediction = model(validation, target, mask) residual = (prediction[mask] - validation[target][mask]).abs() smooth = F.smooth_l1_loss(prediction[mask], validation[target][mask]) rows.append( { "target_modality": target, "mask_ratio": ratio, "mae_standardized": float(residual.mean().item()), "smooth_l1_standardized": float(smooth.item()), "masked_values": int(residual.numel()), } ) return rows def run_shuffled_alignment_reconstruction_probe( train_ids: Sequence[str], val_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], *, device: torch.device, seed: int, ratio: float = 0.2, epochs: int = 25, batch_size: int = 16, learning_rate: float = 1e-3, shuffle_repeats: int = 5, ) -> list[dict[str, float | str]]: """Compare normal reconstruction with cross-modal slots shuffled at eval. The decoder is trained once on aligned training-fold representations. For the control, the two non-target modalities share a random slot permutation within each validation clip; the target stream and target values remain in their original order. Thus the metric isolates how much correctly matched cross-modal slots help the frozen decoder. """ if not train_ids or not val_ids: raise ValueError("reconstruction needs non-empty train and validation sets") if shuffle_repeats < 1: raise ValueError("shuffle_repeats must be positive") torch.manual_seed(seed) generator = torch.Generator(device=device) generator.manual_seed(seed) train = _stack_aligned(train_ids, aligned_by_id, device) validation = _stack_aligned(val_ids, aligned_by_id, device) dimensions = {name: int(train[name].shape[-1]) for name in MODALITIES} model = MaskedCrossModalDecoder(dimensions).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4) rng = np.random.default_rng(seed) model.train() for _ in range(epochs): order = rng.permutation(len(train_ids)) for start in range(0, len(order), batch_size): indices = torch.as_tensor(order[start : start + batch_size], device=device) batch = {name: value.index_select(0, indices) for name, value in train.items()} target_losses = [] for target in MODALITIES: mask = make_block_mask( len(indices), batch[target].shape[1], ratio, device, generator=generator ) prediction = model(batch, target, mask) target_losses.append(F.smooth_l1_loss(prediction[mask], batch[target][mask])) loss = torch.stack(target_losses).mean() optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() model.eval() rows: list[dict[str, float | str]] = [] with torch.no_grad(): for target_index, target in enumerate(MODALITIES): mask = _fixed_block_mask( len(val_ids), validation[target].shape[1], ratio, device, target_index ) aligned_prediction = model(validation, target, mask) aligned_error = (aligned_prediction[mask] - validation[target][mask]).abs().mean() shuffle_errors = [] permutation_rng = np.random.default_rng(seed + 1709 + target_index) slot_count = validation[target].shape[1] for _ in range(shuffle_repeats): permutations = np.stack( [permutation_rng.permutation(slot_count) for _ in val_ids] ) permutation_tensor = torch.as_tensor(permutations, dtype=torch.long, device=device) shuffled = {} for name, values in validation.items(): if name == target: shuffled[name] = values else: gather_indices = permutation_tensor.unsqueeze(-1).expand_as(values) shuffled[name] = values.gather(1, gather_indices) shuffled_prediction = model(shuffled, target, mask) shuffled_error = ( shuffled_prediction[mask] - validation[target][mask] ).abs().mean() shuffle_errors.append(float(shuffled_error.item())) shuffled_mean = float(np.mean(shuffle_errors)) rows.append( { "target_modality": target, "mask_ratio": ratio, "mae_aligned": float(aligned_error.item()), "mae_shuffled_mean": shuffled_mean, "mae_shuffled_std": float(np.std(shuffle_errors, ddof=1)) if shuffle_repeats > 1 else 0.0, "gain_align": shuffled_mean - float(aligned_error.item()), "shuffle_repeats": float(shuffle_repeats), "masked_values": int(mask.sum().item() * dimensions[target]), } ) return rows def _emotion_features( sample_ids: Sequence[str], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], bins: int = 5 ) -> np.ndarray: outputs = [] for sample_id in sample_ids: combined = np.concatenate([aligned_by_id[sample_id][name] for name in MODALITIES], axis=-1) segments = np.array_split(combined, bins, axis=0) outputs.append(np.concatenate([segment.mean(axis=0) for segment in segments])) return np.stack(outputs).astype(np.float32, copy=False) def run_frozen_emotion_probe( train_samples: Sequence[object], val_samples: Sequence[object], aligned_by_id: Mapping[str, Mapping[str, np.ndarray]], ) -> dict[str, float | int]: """Evaluate a regularized five-bin linear probe on frozen aligned features.""" train_ids = [sample.sample_id for sample in train_samples] val_ids = [sample.sample_id for sample in val_samples] x_train = _emotion_features(train_ids, aligned_by_id) x_val = _emotion_features(val_ids, aligned_by_id) y_train = np.asarray([sample.polarity for sample in train_samples], dtype=np.int64) y_val = np.asarray([sample.polarity for sample in val_samples], dtype=np.int64) target_train = np.asarray([sample.sentiment for sample in train_samples], dtype=np.float64) target_val = np.asarray([sample.sentiment for sample in val_samples], dtype=np.float64) scaler = StandardScaler() x_train = scaler.fit_transform(x_train) x_val = scaler.transform(x_val) if np.unique(y_train).size > 1: classifier = LogisticRegression( C=0.1, class_weight="balanced", max_iter=2000, solver="lbfgs", random_state=0 ) else: classifier = DummyClassifier(strategy="most_frequent") classifier.fit(x_train, y_train) predicted_class = classifier.predict(x_val) regressor = Ridge(alpha=10.0) regressor.fit(x_train, target_train) predicted_score = regressor.predict(x_val) if len(target_val) > 1 and np.std(predicted_score) > 0 and np.std(target_val) > 0: pearson = float(np.corrcoef(predicted_score, target_val)[0, 1]) else: pearson = float("nan") return { "accuracy": float(accuracy_score(y_val, predicted_class)), "macro_f1": float(f1_score(y_val, predicted_class, average="macro", zero_division=0)), "mae": float(mean_absolute_error(target_val, predicted_score)), "pearson": pearson, "n_train": len(train_samples), "n_validation": len(val_samples), }