463 lines
19 KiB
Python
463 lines
19 KiB
Python
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),
|
|
}
|