Files
modeling_zhaocui/deep_learning/Q1/q1/experiment_probes.py
T

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