建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,462 @@
|
||||
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),
|
||||
}
|
||||
Reference in New Issue
Block a user