Files

491 lines
23 KiB
Python

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",)
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"}
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="Train the EarlyConcat baseline and its alignment-transfer control")
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" / "followups" / "earlyconcat_standalone"))
args = parser.parse_args()
_run(args)
if __name__ == "__main__":
main()