整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Q2 robustness and Q3 explanation-selection experiments."""
|
||||
@@ -0,0 +1,213 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pickle
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[3]
|
||||
ATTACHMENT2 = ROOT / "E题数据" / "附件2-数据集特征文件"
|
||||
MODALITIES = ("text", "audio", "vision")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Split:
|
||||
x: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
mask: np.ndarray # N x T x 3
|
||||
y_cls: np.ndarray
|
||||
y_reg: np.ndarray
|
||||
ids: list[str]
|
||||
|
||||
@property
|
||||
def n(self) -> int:
|
||||
return len(self.y_cls)
|
||||
|
||||
@property
|
||||
def steps(self) -> int:
|
||||
return int(self.x[0].shape[1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class RobustStats:
|
||||
center: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
scale: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
|
||||
def save(self, path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
np.savez_compressed(
|
||||
path,
|
||||
text_center=self.center[0], text_scale=self.scale[0],
|
||||
audio_center=self.center[1], audio_scale=self.scale[1],
|
||||
vision_center=self.center[2], vision_scale=self.scale[2],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: Path) -> "RobustStats":
|
||||
with np.load(path) as data:
|
||||
return cls(
|
||||
tuple(data[f"{m}_center"].astype(np.float32) for m in MODALITIES),
|
||||
tuple(data[f"{m}_scale"].astype(np.float32) for m in MODALITIES),
|
||||
)
|
||||
|
||||
|
||||
def _unpickle(path: Path) -> dict[str, Any]:
|
||||
with path.open("rb") as stream:
|
||||
return pickle.load(stream, encoding="latin1")
|
||||
|
||||
|
||||
def _ids_and_targets(part: dict[str, Any]) -> tuple[list[str], np.ndarray, np.ndarray]:
|
||||
ids = [str(x) for x in part["id"]]
|
||||
y_cls = np.asarray(part["classification_labels"], dtype=np.int64).reshape(-1)
|
||||
y_reg = np.asarray(part["regression_labels"], dtype=np.float32).reshape(-1)
|
||||
return ids, y_cls, y_reg
|
||||
|
||||
|
||||
def _text_mask(part: dict[str, Any]) -> np.ndarray:
|
||||
tokens = np.asarray(part["text_bert"])
|
||||
if tokens.ndim != 3 or tokens.shape[1] < 2:
|
||||
raise ValueError(f"unexpected text_bert shape: {tokens.shape}")
|
||||
# MOSEI text_bert rows are input_ids, input_mask, segment_ids.
|
||||
return tokens[:, 1, :].astype(bool)
|
||||
|
||||
|
||||
def load_aligned(path: Path | None = None) -> dict[str, Split]:
|
||||
path = path or ATTACHMENT2 / "aligned_50.pkl"
|
||||
raw = _unpickle(path)
|
||||
result: dict[str, Split] = {}
|
||||
for name in ("train", "valid"):
|
||||
part = raw[name]
|
||||
xs = tuple(np.asarray(part[m], dtype=np.float32) for m in MODALITIES)
|
||||
masks = [
|
||||
_text_mask(part),
|
||||
np.any(np.isfinite(xs[1]) & (xs[1] != 0), axis=-1),
|
||||
np.any(np.isfinite(xs[2]) & (xs[2] != 0), axis=-1),
|
||||
]
|
||||
mask = np.stack(masks, axis=-1)
|
||||
ids, y_cls, y_reg = _ids_and_targets(part)
|
||||
if any(x.shape[1] != 50 for x in xs):
|
||||
raise ValueError(f"{name} aligned feature tensors must have 50 slots")
|
||||
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
||||
train_videos = {x.split("$_$", 1)[0] for x in result["train"].ids}
|
||||
valid_videos = {x.split("$_$", 1)[0] for x in result["valid"].ids}
|
||||
overlap = train_videos & valid_videos
|
||||
if overlap:
|
||||
raise ValueError(f"official train/valid split leaks {len(overlap)} source video ids")
|
||||
return result
|
||||
|
||||
|
||||
def _resample_rows_to_50(values: np.ndarray, lengths: list[int] | np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
||||
n, source_steps, dim = values.shape
|
||||
output = np.zeros((n, 50, dim), dtype=np.float32)
|
||||
mask = np.zeros((n, 50), dtype=bool)
|
||||
lengths_arr = np.asarray(lengths, dtype=np.int64).reshape(-1)
|
||||
for i in range(n):
|
||||
length = int(np.clip(lengths_arr[i], 0, source_steps))
|
||||
if length == 0:
|
||||
continue
|
||||
source = np.nan_to_num(values[i, :length], nan=0.0, posinf=0.0, neginf=0.0)
|
||||
observed = np.any(source != 0, axis=-1)
|
||||
for j in range(50):
|
||||
left = int(np.floor(j * length / 50))
|
||||
right = max(left + 1, int(np.ceil((j + 1) * length / 50)))
|
||||
right = min(right, length)
|
||||
use = observed[left:right]
|
||||
if use.any():
|
||||
output[i, j] = source[left:right][use].mean(axis=0)
|
||||
mask[i, j] = True
|
||||
return output, mask
|
||||
|
||||
|
||||
def load_fixed_window(path: Path | None = None) -> dict[str, Split]:
|
||||
"""Build a matched 50-slot equal-window control from the unaligned file."""
|
||||
path = path or ATTACHMENT2 / "unaligned_50.pkl"
|
||||
raw = _unpickle(path)
|
||||
result: dict[str, Split] = {}
|
||||
for name in ("train", "valid"):
|
||||
part = raw[name]
|
||||
text = np.asarray(part["text"], dtype=np.float32)
|
||||
audio, audio_mask = _resample_rows_to_50(part["audio"], part["audio_lengths"])
|
||||
vision, vision_mask = _resample_rows_to_50(part["vision"], part["vision_lengths"])
|
||||
text_mask = _text_mask(part)
|
||||
xs = (text, audio, vision)
|
||||
mask = np.stack((text_mask, audio_mask, vision_mask), axis=-1)
|
||||
ids, y_cls, y_reg = _ids_and_targets(part)
|
||||
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
||||
return result
|
||||
|
||||
|
||||
def fit_robust_stats(split: Split) -> RobustStats:
|
||||
centers: list[np.ndarray] = []
|
||||
scales: list[np.ndarray] = []
|
||||
for modality in range(3):
|
||||
observed = split.mask[:, :, modality].reshape(-1)
|
||||
values = split.x[modality].reshape(-1, split.x[modality].shape[-1])[observed]
|
||||
if not len(values):
|
||||
raise ValueError(f"no observed values for {MODALITIES[modality]}")
|
||||
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
||||
center = np.median(values, axis=0)
|
||||
mad = np.median(np.abs(values - center), axis=0)
|
||||
scale = 1.4826 * mad
|
||||
std = np.std(values, axis=0)
|
||||
scale = np.where(scale > 1e-6, scale, std)
|
||||
scale = np.where(scale > 1e-6, scale, 1.0)
|
||||
centers.append(center.astype(np.float32))
|
||||
scales.append(scale.astype(np.float32))
|
||||
return RobustStats(tuple(centers), tuple(scales))
|
||||
|
||||
|
||||
def apply_robust_stats(split: Split, stats: RobustStats) -> Split:
|
||||
xs: list[np.ndarray] = []
|
||||
for modality in range(3):
|
||||
values = (split.x[modality] - stats.center[modality]) / stats.scale[modality]
|
||||
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
||||
values *= split.mask[:, :, modality, None]
|
||||
xs.append(values.astype(np.float32, copy=False))
|
||||
return Split(tuple(xs), split.mask.copy(), split.y_cls, split.y_reg, split.ids)
|
||||
|
||||
|
||||
def corrupt_masks(
|
||||
base: np.ndarray,
|
||||
ratio: float,
|
||||
modalities: tuple[int, ...],
|
||||
seed: int,
|
||||
) -> np.ndarray:
|
||||
result = base.copy()
|
||||
rng = np.random.default_rng(seed)
|
||||
n, steps, _ = result.shape
|
||||
width = max(1, min(steps, int(round(ratio * steps))))
|
||||
starts = rng.integers(0, steps - width + 1, size=n)
|
||||
for row, start in enumerate(starts.tolist()):
|
||||
result[row, start:start + width, list(modalities)] = False
|
||||
return result
|
||||
|
||||
|
||||
def augment_masks(base: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
||||
result = base.copy()
|
||||
n, steps, _ = result.shape
|
||||
for row in range(n):
|
||||
if rng.random() >= 0.85:
|
||||
continue
|
||||
count = int(rng.integers(1, 4))
|
||||
modalities = rng.choice(3, size=count, replace=False)
|
||||
ratio = float(rng.choice((0.10, 0.20, 0.30)))
|
||||
width = max(1, int(round(ratio * steps)))
|
||||
start = int(rng.integers(0, steps - width + 1))
|
||||
result[row, start:start + width, modalities] = False
|
||||
return result
|
||||
|
||||
|
||||
def shift_audio_vision(split: Split, seed: int, max_shift: int = 10) -> Split:
|
||||
rng = np.random.default_rng(seed)
|
||||
xs = [x.copy() for x in split.x]
|
||||
masks = split.mask.copy()
|
||||
for row in range(split.n):
|
||||
for modality in (1, 2):
|
||||
shift = int(rng.integers(1, max_shift + 1))
|
||||
if rng.random() < 0.5:
|
||||
shift = -shift
|
||||
xs[modality][row] = np.roll(xs[modality][row], shift, axis=0)
|
||||
masks[row, :, modality] = np.roll(masks[row, :, modality], shift)
|
||||
return Split(tuple(xs), masks, split.y_cls, split.y_reg, split.ids)
|
||||
@@ -0,0 +1,29 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
from .train_compare import _plot, _summary, _write_csv
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Rebuild Q2 summary tables from saved validation predictions")
|
||||
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "algorithm_selection"))
|
||||
args = parser.parse_args()
|
||||
output = Path(args.output_dir)
|
||||
with (output / "validation_metrics_by_condition.csv").open(encoding="utf-8-sig", newline="") as stream:
|
||||
rows = list(csv.DictReader(stream))
|
||||
for row in rows:
|
||||
for key in ("missing_rate", "accuracy", "macro_f1", "mae", "pearson", "n_valid"):
|
||||
row[key] = float(row[key])
|
||||
row["seed"] = int(row["seed"])
|
||||
summary = _summary(rows)
|
||||
_write_csv(output / "summary.csv", summary)
|
||||
aligned = [row for row in summary if row["representation"] == "provided_word_aligned_50"]
|
||||
_plot(aligned, rows, output / "missing_rate_comparison.png")
|
||||
print(f"rebuilt summary table and plot from {len(rows)} saved validation rows")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class AlignedFusionModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
kind: str,
|
||||
dims: tuple[int, int, int],
|
||||
steps: int = 50,
|
||||
hidden: int = 128,
|
||||
dropout: float = 0.15,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if kind not in {"concat", "gate", "crossattn"}:
|
||||
raise ValueError(f"unknown model kind: {kind}")
|
||||
self.kind = kind
|
||||
self.hidden = hidden
|
||||
self.projections = nn.ModuleList(
|
||||
nn.Sequential(nn.Linear(size, hidden), nn.GELU(), nn.LayerNorm(hidden))
|
||||
for size in dims
|
||||
)
|
||||
self.position = nn.Parameter(torch.randn(1, steps, hidden) * 0.02)
|
||||
self.modality = nn.Parameter(torch.randn(1, 1, 3, hidden) * 0.02)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
if kind == "concat":
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden * 3 + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
elif kind == "gate":
|
||||
self.gate_score = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.Tanh(), nn.Linear(hidden // 2, 1))
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
else:
|
||||
layer = nn.TransformerEncoderLayer(
|
||||
d_model=hidden,
|
||||
nhead=4,
|
||||
dim_feedforward=hidden * 2,
|
||||
dropout=dropout,
|
||||
activation="gelu",
|
||||
batch_first=True,
|
||||
norm_first=True,
|
||||
)
|
||||
self.cross_encoder = nn.TransformerEncoder(layer, num_layers=2, enable_nested_tensor=False)
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
self.temporal = nn.GRU(
|
||||
input_size=hidden,
|
||||
hidden_size=hidden // 2,
|
||||
num_layers=1,
|
||||
batch_first=True,
|
||||
bidirectional=True,
|
||||
)
|
||||
self.head = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout))
|
||||
self.classifier = nn.Linear(hidden // 2, 3)
|
||||
self.regressor = nn.Linear(hidden // 2, 1)
|
||||
|
||||
def forward(self, xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: torch.Tensor):
|
||||
masks = masks.bool()
|
||||
pos = self.position[:, :masks.shape[1]]
|
||||
encoded = []
|
||||
for modality, (projection, x) in enumerate(zip(self.projections, xs)):
|
||||
token = projection(x)
|
||||
token = token + pos + self.modality[:, :, modality, :]
|
||||
token = token * masks[:, :, modality, None]
|
||||
encoded.append(token)
|
||||
stack = torch.stack(encoded, dim=2) # B x T x M x D
|
||||
availability = masks.to(stack.dtype)
|
||||
gate_weights = None
|
||||
|
||||
if self.kind == "concat":
|
||||
fused = self.fusion(torch.cat((stack.flatten(2), availability), dim=-1))
|
||||
elif self.kind == "gate":
|
||||
scores = self.gate_score(stack).squeeze(-1)
|
||||
scores = scores.masked_fill(~masks, -1e4)
|
||||
gate_weights = torch.softmax(scores, dim=-1) * availability
|
||||
gate_weights = gate_weights / gate_weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
|
||||
weighted = (stack * gate_weights[..., None]).sum(dim=2)
|
||||
fused = self.fusion(torch.cat((weighted, availability), dim=-1))
|
||||
else:
|
||||
batch, steps, modalities, hidden = stack.shape
|
||||
flat = stack.reshape(batch, steps * modalities, hidden)
|
||||
valid = masks.reshape(batch, steps * modalities).clone()
|
||||
empty = ~valid.any(dim=1)
|
||||
if empty.any():
|
||||
valid[empty, 0] = True
|
||||
flat[empty, 0] = 0.0
|
||||
attended = self.cross_encoder(flat, src_key_padding_mask=~valid)
|
||||
attended = attended.reshape(batch, steps, modalities, hidden)
|
||||
observed_count = availability.sum(dim=2, keepdim=True)
|
||||
pooled = (attended * availability[..., None]).sum(dim=2) / observed_count.clamp_min(1.0)
|
||||
fused = self.fusion(torch.cat((pooled, availability), dim=-1))
|
||||
|
||||
temporal, _ = self.temporal(self.dropout(fused))
|
||||
time_weight = masks.any(dim=-1).to(temporal.dtype)
|
||||
empty_time = time_weight.sum(dim=1, keepdim=True) <= 0
|
||||
if empty_time.any():
|
||||
time_weight[empty_time.squeeze(1), 0] = 1.0
|
||||
pooled = (temporal * time_weight[..., None]).sum(dim=1) / time_weight.sum(dim=1, keepdim=True).clamp_min(1.0)
|
||||
hidden = self.head(pooled)
|
||||
logits = self.classifier(hidden)
|
||||
intensity = 3.0 * torch.tanh(self.regressor(hidden).squeeze(-1))
|
||||
return {"logits": logits, "intensity": intensity, "gate": gate_weights}
|
||||
@@ -0,0 +1,490 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import shutil
|
||||
import time
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
ROOT,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
apply_robust_stats,
|
||||
augment_masks,
|
||||
corrupt_masks,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
load_fixed_window,
|
||||
shift_audio_vision,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
|
||||
|
||||
PATTERNS = {
|
||||
"text": (0,),
|
||||
"audio": (1,),
|
||||
"vision": (2,),
|
||||
"audio_vision": (1, 2),
|
||||
"all_modalities": (0, 1, 2),
|
||||
}
|
||||
KINDS = ("concat", "gate", "crossattn")
|
||||
|
||||
|
||||
def seed_everything(seed: int) -> None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
def _tensor_split(split: Split, device: torch.device) -> tuple[tuple[torch.Tensor, ...], torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in split.x)
|
||||
mask = torch.as_tensor(split.mask, dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(split.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(split.y_reg, dtype=torch.float32, device=device)
|
||||
return xs, mask, y_cls, y_reg
|
||||
|
||||
|
||||
def _loss(output: dict[str, torch.Tensor], y_cls: torch.Tensor, y_reg: torch.Tensor) -> torch.Tensor:
|
||||
class_loss = F.cross_entropy(output["logits"], y_cls)
|
||||
intensity_loss = F.smooth_l1_loss(output["intensity"] / 3.0, y_reg / 3.0)
|
||||
return class_loss + 0.5 * intensity_loss
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _score_arrays(
|
||||
model: AlignedFusionModel,
|
||||
split: Split,
|
||||
mask: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int = 128,
|
||||
) -> tuple[dict[str, float], dict[str, np.ndarray]]:
|
||||
model.eval()
|
||||
predictions: dict[str, list[np.ndarray]] = {"logits": [], "intensity": []}
|
||||
xs = split.x
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xb = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs)
|
||||
mb = torch.as_tensor(mask[start:end], dtype=torch.bool, device=device)
|
||||
output = model(xb, mb)
|
||||
predictions["logits"].append(output["logits"].float().cpu().numpy())
|
||||
predictions["intensity"].append(output["intensity"].float().cpu().numpy())
|
||||
logits = np.concatenate(predictions["logits"], axis=0)
|
||||
intensity = np.clip(np.concatenate(predictions["intensity"], axis=0), -3.0, 3.0)
|
||||
pred_cls = logits.argmax(axis=-1)
|
||||
pearson = _pearson(split.y_reg, intensity)
|
||||
metrics = {
|
||||
"accuracy": float(accuracy_score(split.y_cls, pred_cls)),
|
||||
"macro_f1": float(f1_score(split.y_cls, pred_cls, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(split.y_reg, intensity)),
|
||||
"pearson": pearson,
|
||||
}
|
||||
return metrics, {"logits": logits, "intensity": intensity, "class": pred_cls}
|
||||
|
||||
|
||||
def _pearson(y: np.ndarray, pred: np.ndarray) -> float:
|
||||
a = np.asarray(y, dtype=np.float64)
|
||||
b = np.asarray(pred, dtype=np.float64)
|
||||
if a.std() < 1e-12 or b.std() < 1e-12:
|
||||
return 0.0
|
||||
return float(np.corrcoef(a, b)[0, 1])
|
||||
|
||||
|
||||
def _validation_loss(model: AlignedFusionModel, valid: Split, device: torch.device, batch_size: int) -> float:
|
||||
model.eval()
|
||||
xs, masks, y_cls, y_reg = _tensor_split(valid, device)
|
||||
losses: list[float] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, valid.n, batch_size):
|
||||
idx = slice(start, min(start + batch_size, valid.n))
|
||||
output = model(tuple(x[idx] for x in xs), masks[idx])
|
||||
losses.append(float(_loss(output, y_cls[idx], y_reg[idx]).item()))
|
||||
return float(np.average(losses, weights=[min(batch_size, valid.n - i) for i in range(0, valid.n, batch_size)]))
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if not rows:
|
||||
return
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _train_one(
|
||||
kind: str,
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
) -> tuple[AlignedFusionModel, int, list[dict[str, float]]]:
|
||||
seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
model = AlignedFusionModel(kind, dims=dims).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-4, weight_decay=1e-4)
|
||||
train_tensors = _tensor_split(train, device)
|
||||
xs, base_masks, y_cls, y_reg = train_tensors
|
||||
rng = np.random.default_rng(seed + 809)
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale_epochs = 0
|
||||
history: list[dict[str, float]] = []
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
order = rng.permutation(train.n)
|
||||
batch_losses: list[float] = []
|
||||
for start in range(0, train.n, batch_size):
|
||||
ids_np = order[start:start + batch_size]
|
||||
ids = torch.as_tensor(ids_np, dtype=torch.long, device=device)
|
||||
masks_np = augment_masks(train.mask[ids_np], rng)
|
||||
masks = torch.as_tensor(masks_np, dtype=torch.bool, device=device)
|
||||
output = model(tuple(x.index_select(0, ids) for x in xs), masks)
|
||||
loss = _loss(output, y_cls.index_select(0, ids), y_reg.index_select(0, ids))
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
batch_losses.append(float(loss.detach().item()))
|
||||
valid_loss = _validation_loss(model, valid, device, batch_size)
|
||||
row = {"epoch": float(epoch), "train_loss": float(np.mean(batch_losses)), "valid_clean_loss": valid_loss}
|
||||
history.append(row)
|
||||
print(f"[{kind}] epoch={epoch:02d} train={row['train_loss']:.4f} valid={valid_loss:.4f}", flush=True)
|
||||
if valid_loss < best_loss - 1e-4:
|
||||
best_loss = valid_loss
|
||||
best_epoch = epoch
|
||||
stale_epochs = 0
|
||||
torch.save({"kind": kind, "dims": dims, "state_dict": model.state_dict(), "seed": seed, "best_epoch": epoch}, checkpoint_path)
|
||||
else:
|
||||
stale_epochs += 1
|
||||
if stale_epochs >= patience:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
_write_csv(output_dir / "training_history.csv", history)
|
||||
return model, best_epoch, history
|
||||
|
||||
|
||||
def _conditions(valid: Split, seed: int) -> list[tuple[str, float, np.ndarray]]:
|
||||
result = [("clean", 0.0, valid.mask.copy())]
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
for pattern_id, (pattern, mods) in enumerate(PATTERNS.items()):
|
||||
result.append((pattern, rate, corrupt_masks(valid.mask, rate, mods, seed + pattern_id * 101 + int(rate * 1000))))
|
||||
return result
|
||||
|
||||
|
||||
def _eval_conditions(
|
||||
model: AlignedFusionModel,
|
||||
valid: Split,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
seed_run: int,
|
||||
method: str,
|
||||
representation: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
rows = []
|
||||
for condition, rate, masks in _conditions(valid, seed):
|
||||
metrics, _ = _score_arrays(model, valid, masks, device)
|
||||
rows.append({"method": method, "representation": representation, "seed": seed_run, "condition": condition,
|
||||
"missing_rate": rate, "n_valid": valid.n, **metrics})
|
||||
print(f"[{method}/{representation}] {condition:14s} rate={rate:.1f} "
|
||||
f"F1={metrics['macro_f1']:.3f} MAE={metrics['mae']:.3f} "
|
||||
f"P={metrics['pearson']:.3f}", flush=True)
|
||||
return rows
|
||||
|
||||
|
||||
def _summary(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
groups = list(dict.fromkeys((row["method"], row["representation"]) for row in rows))
|
||||
summary: list[dict[str, Any]] = []
|
||||
for method, representation in groups:
|
||||
matching = [r for r in rows if r["method"] == method and r["representation"] == representation]
|
||||
local = [r for r in matching if r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
clean = [r for r in matching if r["condition"] == "clean"]
|
||||
seeds = sorted({int(r.get("seed", 0)) for r in matching})
|
||||
|
||||
def per_seed_mean(selected: list[dict[str, Any]], metric: str) -> list[float]:
|
||||
return [float(np.mean([r[metric] for r in selected if int(r.get("seed", 0)) == seed]))
|
||||
for seed in seeds if any(int(r.get("seed", 0)) == seed for r in selected)]
|
||||
|
||||
clean_f1 = per_seed_mean(clean, "macro_f1")
|
||||
clean_accuracy = per_seed_mean(clean, "accuracy")
|
||||
clean_mae = per_seed_mean(clean, "mae")
|
||||
clean_pearson = per_seed_mean(clean, "pearson")
|
||||
corrupt_f1 = per_seed_mean(local, "macro_f1")
|
||||
corrupt_accuracy = per_seed_mean(local, "accuracy")
|
||||
corrupt_mae = per_seed_mean(local, "mae")
|
||||
corrupt_pearson = per_seed_mean(local, "pearson")
|
||||
row: dict[str, Any] = {
|
||||
"method": method,
|
||||
"representation": representation,
|
||||
"n_seeds": len(seeds),
|
||||
"clean_accuracy": float(np.mean(clean_accuracy)),
|
||||
"clean_accuracy_sd": float(np.std(clean_accuracy, ddof=1)) if len(clean_accuracy) > 1 else 0.0,
|
||||
"clean_macro_f1": float(np.mean(clean_f1)),
|
||||
"clean_macro_f1_sd": float(np.std(clean_f1, ddof=1)) if len(clean_f1) > 1 else 0.0,
|
||||
"clean_mae": float(np.mean(clean_mae)),
|
||||
"clean_mae_sd": float(np.std(clean_mae, ddof=1)) if len(clean_mae) > 1 else 0.0,
|
||||
"clean_pearson": float(np.mean(clean_pearson)),
|
||||
"clean_pearson_sd": float(np.std(clean_pearson, ddof=1)) if len(clean_pearson) > 1 else 0.0,
|
||||
"corrupt_accuracy_mean": float(np.mean(corrupt_accuracy)),
|
||||
"corrupt_accuracy_sd": float(np.std(corrupt_accuracy, ddof=1)) if len(corrupt_accuracy) > 1 else 0.0,
|
||||
"corrupt_macro_f1_mean": float(np.mean(corrupt_f1)),
|
||||
"corrupt_macro_f1_sd": float(np.std(corrupt_f1, ddof=1)) if len(corrupt_f1) > 1 else 0.0,
|
||||
"corrupt_macro_f1_worst": float(np.min([r["macro_f1"] for r in local])),
|
||||
"corrupt_mae_mean": float(np.mean(corrupt_mae)),
|
||||
"corrupt_mae_sd": float(np.std(corrupt_mae, ddof=1)) if len(corrupt_mae) > 1 else 0.0,
|
||||
"corrupt_pearson_mean": float(np.mean(corrupt_pearson)),
|
||||
"corrupt_pearson_sd": float(np.std(corrupt_pearson, ddof=1)) if len(corrupt_pearson) > 1 else 0.0,
|
||||
}
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
at_rate = [r for r in local if r["missing_rate"] == rate]
|
||||
f1_by_seed = per_seed_mean(at_rate, "macro_f1")
|
||||
accuracy_by_seed = per_seed_mean(at_rate, "accuracy")
|
||||
mae_by_seed = per_seed_mean(at_rate, "mae")
|
||||
row[f"f1_rate_{int(rate * 100)}"] = float(np.mean(f1_by_seed))
|
||||
row[f"accuracy_rate_{int(rate * 100)}"] = float(np.mean(accuracy_by_seed))
|
||||
row[f"mae_rate_{int(rate * 100)}"] = float(np.mean(mae_by_seed))
|
||||
summary.append(row)
|
||||
for row in summary:
|
||||
row["pareto_nondominated"] = not any(
|
||||
other is not row and other["representation"] == row["representation"]
|
||||
and other["corrupt_macro_f1_mean"] >= row["corrupt_macro_f1_mean"]
|
||||
and other["corrupt_mae_mean"] <= row["corrupt_mae_mean"]
|
||||
and other["corrupt_pearson_mean"] >= row["corrupt_pearson_mean"]
|
||||
and (
|
||||
other["corrupt_macro_f1_mean"] > row["corrupt_macro_f1_mean"]
|
||||
or other["corrupt_mae_mean"] < row["corrupt_mae_mean"]
|
||||
or other["corrupt_pearson_mean"] > row["corrupt_pearson_mean"]
|
||||
)
|
||||
for other in summary
|
||||
)
|
||||
return summary
|
||||
|
||||
|
||||
def _plot(summary: list[dict[str, Any]], rows: list[dict[str, Any]], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
colors = {"concat": "#4e79a7", "gate": "#f28e2b", "crossattn": "#59a14f"}
|
||||
fig, axes = plt.subplots(1, 2, figsize=(11, 4.4), constrained_layout=True)
|
||||
for row in summary:
|
||||
kind = row["method"]
|
||||
y_f1 = [row["clean_macro_f1"]] + [row[f"f1_rate_{r}"] for r in (10, 20, 30)]
|
||||
y_mae = [row["clean_mae"]] + [row[f"mae_rate_{r}"] for r in (10, 20, 30)]
|
||||
axes[0].plot([0, 10, 20, 30], y_f1, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[1].plot([0, 10, 20, 30], y_mae, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[0].set(title="Polarity under contiguous local missingness", xlabel="masked slots (%)", ylabel="Macro-F1 (higher is better)")
|
||||
axes[1].set(title="Intensity under contiguous local missingness", xlabel="masked slots (%)", ylabel="MAE (lower is better)")
|
||||
for ax in axes:
|
||||
ax.grid(alpha=0.25)
|
||||
ax.legend(frameon=False)
|
||||
fig.savefig(path, dpi=180)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for block in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _run(args: argparse.Namespace) -> None:
|
||||
seed_everything(args.seeds[0])
|
||||
if args.device == "auto":
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
else:
|
||||
device = torch.device(args.device)
|
||||
torch.set_num_threads(args.threads)
|
||||
output = Path(args.output_dir)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
aligned_raw = load_aligned()
|
||||
stats = fit_robust_stats(aligned_raw["train"])
|
||||
stats.save(output / "aligned_robust_stats.npz")
|
||||
aligned = {k: apply_robust_stats(v, stats) for k, v in aligned_raw.items()}
|
||||
audit = {
|
||||
"source": str(ATTACHMENT2 / "aligned_50.pkl"),
|
||||
"train_samples": aligned["train"].n,
|
||||
"valid_samples": aligned["valid"].n,
|
||||
"train_classes": np.bincount(aligned["train"].y_cls, minlength=3).tolist(),
|
||||
"valid_classes": np.bincount(aligned["valid"].y_cls, minlength=3).tolist(),
|
||||
"mean_observed_slots": {
|
||||
MODALITIES[m]: float(aligned["train"].mask[:, :, m].sum(axis=1).mean()) for m in range(3)
|
||||
},
|
||||
"train_valid_video_overlap": 0,
|
||||
}
|
||||
with (output / "data_audit.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(audit, stream, ensure_ascii=False, indent=2)
|
||||
print(f"device={device}; train={audit['train_samples']}; valid={audit['valid_samples']}; audit={audit}", flush=True)
|
||||
|
||||
metric_rows: list[dict[str, Any]] = []
|
||||
best_epochs: dict[str, int] = {}
|
||||
for kind in KINDS:
|
||||
for seed in args.seeds:
|
||||
seed_dir = output / "models" / "aligned" / kind / f"seed_{seed}"
|
||||
model, best_epoch, _ = _train_one(
|
||||
kind, aligned["train"], aligned["valid"], seed_dir,
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"{kind}_seed_{seed}"] = best_epoch
|
||||
metric_rows.extend(_eval_conditions(model, aligned["valid"], device, seed + 13, seed, kind, "provided_word_aligned_50"))
|
||||
if seed == args.seeds[0]:
|
||||
shutil.copy2(seed_dir / "model_best.pt", output / "models" / "aligned" / kind / "model_best.pt")
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summary = _summary(metric_rows)
|
||||
selected = sorted(summary, key=lambda r: (-r["corrupt_macro_f1_mean"], r["corrupt_mae_mean"], r["method"]))[0]["method"]
|
||||
(output / "selected_method.txt").write_text(
|
||||
f"Macro-F1-first validation selection: {selected}. See summary.csv for the full multi-metric tradeoff.\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Matched audio/vision temporal-shift control for the selected architecture and every seed.
|
||||
for seed in args.seeds:
|
||||
aligned_payload = torch.load(output / "models" / "aligned" / selected / f"seed_{seed}" / "model_best.pt",
|
||||
map_location=device, weights_only=False)
|
||||
aligned_model = AlignedFusionModel(selected, tuple(aligned_payload["dims"])).to(device)
|
||||
aligned_model.load_state_dict(aligned_payload["state_dict"])
|
||||
shifted = shift_audio_vision(aligned["valid"], seed=seed + 2026, max_shift=10)
|
||||
shift_metrics, _ = _score_arrays(aligned_model, shifted, shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "provided_word_aligned_50", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": shifted.n, **shift_metrics})
|
||||
del aligned_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Same selected fusion architecture, but equal-window audio/vision pooling of the unaligned source.
|
||||
print(f"selected_by_corrupt_macro_f1={selected}; starting fixed-window alignment control", flush=True)
|
||||
fixed_raw = load_fixed_window()
|
||||
fixed_stats = fit_robust_stats(fixed_raw["train"])
|
||||
fixed_stats.save(output / "fixed_window_robust_stats.npz")
|
||||
fixed = {k: apply_robust_stats(v, fixed_stats) for k, v in fixed_raw.items()}
|
||||
for seed in args.seeds:
|
||||
fixed_model, fixed_epoch, _ = _train_one(
|
||||
selected, fixed["train"], fixed["valid"], output / "models" / "fixed_window" / selected / f"seed_{seed}",
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"fixed_window_{selected}_seed_{seed}"] = fixed_epoch
|
||||
metric_rows.extend(_eval_conditions(fixed_model, fixed["valid"], device, seed + 13, seed, selected,
|
||||
"equal_window_resampled_unaligned"))
|
||||
fixed_shifted = shift_audio_vision(fixed["valid"], seed=seed + 2026, max_shift=10)
|
||||
fixed_shift_metrics, _ = _score_arrays(fixed_model, fixed_shifted, fixed_shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "equal_window_resampled_unaligned", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": fixed_shifted.n, **fixed_shift_metrics})
|
||||
del fixed_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_summary = _summary(metric_rows)
|
||||
_write_csv(output / "validation_metrics_by_condition.csv", metric_rows)
|
||||
_write_csv(output / "summary.csv", all_summary)
|
||||
aligned_summary = [r for r in all_summary if r["representation"] == "provided_word_aligned_50"]
|
||||
_plot(aligned_summary, metric_rows, output / "missing_rate_comparison.png")
|
||||
alignment_rows = []
|
||||
for rep in ("provided_word_aligned_50", "equal_window_resampled_unaligned"):
|
||||
for condition in ("clean", "audio_vision_shifted_1_to_10_slots"):
|
||||
match = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] == condition]
|
||||
if match:
|
||||
row = {"method": selected, "representation": rep, "condition": condition,
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(match)}
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
values = [r[metric] for r in match]
|
||||
row[metric] = float(np.mean(values))
|
||||
row[f"{metric}_sd"] = float(np.std(values, ddof=1)) if len(values) > 1 else 0.0
|
||||
alignment_rows.append(row)
|
||||
corrupt = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
if corrupt:
|
||||
per_seed = []
|
||||
for seed in args.seeds:
|
||||
local = [r for r in corrupt if int(r["seed"]) == seed]
|
||||
if local:
|
||||
per_seed.append({metric: float(np.mean([r[metric] for r in local])) for metric in
|
||||
("accuracy", "macro_f1", "mae", "pearson")})
|
||||
alignment_rows.append({
|
||||
"method": selected, "representation": rep, "condition": "all_local_corruption_mean",
|
||||
"missing_rate": float(np.mean([r["missing_rate"] for r in corrupt])),
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(per_seed),
|
||||
**{metric: float(np.mean([r[metric] for r in per_seed])) for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
**{f"{metric}_sd": float(np.std([r[metric] for r in per_seed], ddof=1)) if len(per_seed) > 1 else 0.0
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
})
|
||||
_write_csv(output / "alignment_transfer_ablation.csv", alignment_rows)
|
||||
|
||||
source_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
manifest = {
|
||||
"source_feature": str(source_path),
|
||||
"source_sha256": _sha256(source_path),
|
||||
"device": str(device),
|
||||
"cuda_name": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"seeds": args.seeds,
|
||||
"epochs_max": args.epochs,
|
||||
"patience": args.patience,
|
||||
"batch_size": args.batch_size,
|
||||
"best_epochs": best_epochs,
|
||||
"selected_macro_f1_first": selected,
|
||||
"selection_policy": "report Macro-F1, MAE, and Pearson separately; selected model maximizes mean validation Macro-F1 across 15 contiguous corruption conditions, then uses MAE and lexical model name only as tie-breaks",
|
||||
"models": list(KINDS),
|
||||
"corruption_rates": [0.10, 0.20, 0.30],
|
||||
"corruption_patterns": list(PATTERNS),
|
||||
"feature_scaling": "training split median/MAD; fallback to standard deviation for zero-MAD dimensions",
|
||||
"test_labels_used": False,
|
||||
"alignment_transfer_limit": "The official aligned_50 data use a 50-slot wordpiece sequence with no per-slot seconds or stored Q1 B1 time_bounds. The fixed-window comparison is a downstream alignment control, not a re-run of Q1 B1 on the full dataset.",
|
||||
"python": __import__("sys").version,
|
||||
"torch": torch.__version__,
|
||||
"numpy": np.__version__,
|
||||
"created_unix": time.time(),
|
||||
}
|
||||
with (output / "run_manifest.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(manifest, stream, ensure_ascii=False, indent=2)
|
||||
print(f"saved selection artifacts to {output}; selected={selected}; seeds={args.seeds}", flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Q2 local-missingness model and alignment transfer selection")
|
||||
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
|
||||
parser.add_argument("--epochs", type=int, default=32)
|
||||
parser.add_argument("--patience", type=int, default=6)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--threads", type=int, default=4)
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "algorithm_selection"))
|
||||
args = parser.parse_args()
|
||||
_run(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user