491 lines
23 KiB
Python
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()
|