提交其余项目实验变更
This commit is contained in:
@@ -0,0 +1,787 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
apply_robust_stats,
|
||||
augment_masks,
|
||||
corrupt_masks,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
from .mofe import EXPERT_NAMES, SUBSETS, MixtureOfFusionExperts
|
||||
from .train_compare import PATTERNS, _loss, _pearson, _train_one, seed_everything
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
REFERENCE_OUTPUT = ROOT / "outputs" / "mofe_7experts"
|
||||
DEFAULT_OUTPUT = ROOT / "outputs" / "mofe_7experts"
|
||||
EARLYCONCAT = "B0_early_concat"
|
||||
MOFE7_MLP = "B5_mofe_mlp"
|
||||
SEEDS = (42, 3407, 2026)
|
||||
RATES = (0.10, 0.20, 0.30)
|
||||
HIDDEN = 128
|
||||
LATENT_DIM = 64
|
||||
MODEL_CONFIG: dict[str, Any] = {
|
||||
"router": "mlp",
|
||||
"expert_names": EXPERT_NAMES,
|
||||
"availability_mode": "hard",
|
||||
}
|
||||
SUMMARY_METRICS = (
|
||||
"corrupt_macro_f1",
|
||||
"worst_condition_macro_f1",
|
||||
"text_30_macro_f1",
|
||||
"corrupt_mae",
|
||||
"corrupt_pearson",
|
||||
)
|
||||
|
||||
|
||||
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 _read_csv(path: Path) -> list[dict[str, str]]:
|
||||
if not path.exists():
|
||||
return []
|
||||
with path.open("r", newline="", encoding="utf-8-sig") as stream:
|
||||
return list(csv.DictReader(stream))
|
||||
|
||||
|
||||
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 _device_for(name: str) -> torch.device:
|
||||
if name == "auto":
|
||||
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
return torch.device(name)
|
||||
|
||||
|
||||
def _conditions(valid: Split, seed: int) -> list[tuple[str, float, np.ndarray]]:
|
||||
rows = [("clean", 0.0, valid.mask.copy())]
|
||||
for rate in RATES:
|
||||
for pattern_idx, (pattern, modalities) in enumerate(PATTERNS.items()):
|
||||
masks = corrupt_masks(
|
||||
valid.mask,
|
||||
rate,
|
||||
modalities,
|
||||
seed + 13 + pattern_idx * 101 + int(rate * 1000),
|
||||
)
|
||||
rows.append((f"{pattern}_{int(rate * 100)}", rate, masks))
|
||||
return rows
|
||||
|
||||
|
||||
def _metric_dict(
|
||||
y_cls: np.ndarray,
|
||||
y_reg: np.ndarray,
|
||||
logits: np.ndarray,
|
||||
intensity: np.ndarray,
|
||||
) -> dict[str, float]:
|
||||
predicted_class = np.asarray(logits).argmax(axis=-1)
|
||||
predicted_intensity = np.clip(np.asarray(intensity).reshape(-1), -3.0, 3.0)
|
||||
return {
|
||||
"accuracy": float(accuracy_score(y_cls, predicted_class)),
|
||||
"macro_f1": float(f1_score(y_cls, predicted_class, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(y_reg, predicted_intensity)),
|
||||
"pearson": _pearson(y_reg, predicted_intensity),
|
||||
}
|
||||
|
||||
|
||||
def _validation_loss(model: nn.Module, valid: Split, device: torch.device, batch_size: int) -> float:
|
||||
model.eval()
|
||||
values: list[float] = []
|
||||
weights: list[int] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, valid.n, batch_size):
|
||||
end = min(start + batch_size, valid.n)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in valid.x)
|
||||
masks = torch.as_tensor(valid.mask[start:end], dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(valid.y_cls[start:end], dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(valid.y_reg[start:end], dtype=torch.float32, device=device)
|
||||
values.append(float(_loss(model(xs, masks), y_cls, y_reg).item()))
|
||||
weights.append(end - start)
|
||||
return float(np.average(values, weights=weights))
|
||||
|
||||
|
||||
def _train_mofe(
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
reuse_checkpoint: bool,
|
||||
) -> tuple[MixtureOfFusionExperts, int, list[dict[str, Any]]]:
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
history_path = output_dir / "training_history.csv"
|
||||
if reuse_checkpoint and checkpoint_path.exists():
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if saved.get("config") != MODEL_CONFIG or tuple(saved.get("dims", ())) != dims or int(saved.get("seed", -1)) != seed:
|
||||
raise ValueError(f"cached MoFE checkpoint does not match the selected configuration: {checkpoint_path}")
|
||||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
history = [
|
||||
{"method": MOFE7_MLP, "seed": seed, **{key: float(value) for key, value in row.items() if key in {"epoch", "train_loss", "valid_clean_loss"}}}
|
||||
for row in _read_csv(history_path)
|
||||
]
|
||||
return model.eval(), int(saved.get("best_epoch", 0)), history
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
seed_everything(seed)
|
||||
model = MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-4, weight_decay=1e-4)
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in train.x)
|
||||
base_masks = train.mask
|
||||
y_cls = torch.as_tensor(train.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(train.y_reg, dtype=torch.float32, device=device)
|
||||
rng = np.random.default_rng(seed + 809)
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale_epochs = 0
|
||||
history: list[dict[str, Any]] = []
|
||||
|
||||
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(base_masks[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 = {
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"epoch": epoch,
|
||||
"train_loss": float(np.mean(batch_losses)),
|
||||
"valid_clean_loss": valid_loss,
|
||||
}
|
||||
history.append(row)
|
||||
print(f"[MoFE-7 MLP] seed={seed} 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({
|
||||
"method": MOFE7_MLP,
|
||||
"config": MODEL_CONFIG,
|
||||
"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(history_path, history)
|
||||
return model, best_epoch, history
|
||||
|
||||
|
||||
def _load_or_train_concat(
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
reuse_checkpoint: bool,
|
||||
) -> tuple[AlignedFusionModel, int, list[dict[str, Any]]]:
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
if reuse_checkpoint and checkpoint_path.exists():
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
if saved.get("kind") != "concat" or tuple(saved.get("dims", ())) != dims or int(saved.get("seed", -1)) != seed:
|
||||
raise ValueError(f"cached EarlyConcat checkpoint does not match: {checkpoint_path}")
|
||||
model = AlignedFusionModel("concat", dims=dims).to(device)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
history = [
|
||||
{"method": EARLYCONCAT, "seed": seed, **{key: float(value) for key, value in row.items() if key in {"epoch", "train_loss", "valid_clean_loss"}}}
|
||||
for row in _read_csv(output_dir / "training_history.csv")
|
||||
]
|
||||
return model.eval(), int(saved.get("best_epoch", 0)), history
|
||||
|
||||
model, best_epoch, history = _train_one(
|
||||
"concat", train, valid, output_dir, device, seed, epochs, patience, batch_size
|
||||
)
|
||||
rows = [{"method": EARLYCONCAT, "seed": seed, **row} for row in history]
|
||||
return model.eval(), best_epoch, rows
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _predict(
|
||||
model: nn.Module,
|
||||
split: Split,
|
||||
masks: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int,
|
||||
force_expert: str | None = None,
|
||||
) -> dict[str, np.ndarray]:
|
||||
fields = ["logits", "intensity"]
|
||||
if isinstance(model, MixtureOfFusionExperts):
|
||||
fields.extend(("alpha", "utility", "availability", "fallback"))
|
||||
chunks: dict[str, list[np.ndarray]] = {name: [] for name in fields}
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xs = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in split.x)
|
||||
mask_batch = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device)
|
||||
output = model(xs, mask_batch, force_expert=force_expert) if isinstance(model, MixtureOfFusionExperts) else model(xs, mask_batch)
|
||||
for name in fields:
|
||||
value = output[name]
|
||||
chunks[name].append(value.float().cpu().numpy())
|
||||
result = {name: np.concatenate(values, axis=0) for name, values in chunks.items()}
|
||||
result["intensity"] = np.clip(result["intensity"].reshape(-1), -3.0, 3.0)
|
||||
return result
|
||||
|
||||
|
||||
def _condition_row(
|
||||
method: str,
|
||||
seed: int,
|
||||
condition: str,
|
||||
rate: float,
|
||||
split: Split,
|
||||
prediction: dict[str, np.ndarray],
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"method": method,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"missing_rate": rate,
|
||||
"n_valid": split.n,
|
||||
**_metric_dict(split.y_cls, split.y_reg, prediction["logits"], prediction["intensity"]),
|
||||
}
|
||||
|
||||
|
||||
def _diagnostics(
|
||||
seed: int,
|
||||
condition: str,
|
||||
masks: np.ndarray,
|
||||
prediction: dict[str, np.ndarray],
|
||||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
alpha = prediction["alpha"]
|
||||
availability = prediction["availability"].astype(bool)
|
||||
active = availability.any(axis=-1)
|
||||
active_alpha = alpha[active]
|
||||
if active_alpha.size:
|
||||
means = active_alpha.mean(axis=0)
|
||||
entropy = -(active_alpha * np.log(np.maximum(active_alpha, 1e-12))).sum(axis=-1) / np.log(len(EXPERT_NAMES))
|
||||
high_weight = (active_alpha.max(axis=-1) > 0.8).mean()
|
||||
else:
|
||||
means = np.zeros(len(EXPERT_NAMES), dtype=np.float64)
|
||||
entropy = np.zeros(0, dtype=np.float64)
|
||||
high_weight = 0.0
|
||||
route_row: dict[str, Any] = {
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"active_position_fraction": float(active.mean()),
|
||||
"fallback_position_fraction": float((~active).mean()),
|
||||
"normalized_router_entropy": float(entropy.mean()) if entropy.size else 0.0,
|
||||
"fraction_active_positions_max_weight_over_0p8": float(high_weight),
|
||||
}
|
||||
for index, name in enumerate(EXPERT_NAMES):
|
||||
route_row[f"alpha_{name}_mean"] = float(means[index])
|
||||
utility = prediction["utility"]
|
||||
utility_row: dict[str, Any] = {"method": MOFE7_MLP, "seed": seed, "condition": condition}
|
||||
for modality, name in enumerate(MODALITIES):
|
||||
observed = masks[..., modality]
|
||||
utility_row[f"utility_{name}_mean"] = float(utility[..., modality][observed].mean()) if observed.any() else 0.0
|
||||
return route_row, utility_row
|
||||
|
||||
|
||||
def _summary_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
summaries: list[dict[str, Any]] = []
|
||||
for method in (EARLYCONCAT, MOFE7_MLP):
|
||||
matching = [row for row in rows if row["method"] == method]
|
||||
seeds = sorted({int(row["seed"]) for row in matching})
|
||||
conditions = list(dict.fromkeys(row["condition"] for row in matching))
|
||||
by_seed_condition = {(int(row["seed"]), row["condition"]): row for row in matching}
|
||||
clean = [by_seed_condition[(seed, "clean")] for seed in seeds]
|
||||
corrupt_conditions = [condition for condition in conditions if condition != "clean"]
|
||||
corrupt_by_seed = {
|
||||
seed: [by_seed_condition[(seed, condition)] for condition in corrupt_conditions]
|
||||
for seed in seeds
|
||||
}
|
||||
condition_f1 = {
|
||||
condition: float(np.mean([by_seed_condition[(seed, condition)]["macro_f1"] for seed in seeds]))
|
||||
for condition in corrupt_conditions
|
||||
}
|
||||
worst_condition = min(condition_f1, key=condition_f1.get)
|
||||
row: dict[str, Any] = {"method": method, "n_seeds": len(seeds), "worst_condition": worst_condition}
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
clean_values = [float(item[metric]) for item in clean]
|
||||
corrupt_values = [float(np.mean([item[metric] for item in corrupt_by_seed[seed]])) for seed in seeds]
|
||||
row[f"clean_{metric}"] = float(np.mean(clean_values))
|
||||
row[f"clean_{metric}_sd"] = float(np.std(clean_values, ddof=1)) if len(clean_values) > 1 else 0.0
|
||||
row[f"corrupt_{metric}_mean"] = float(np.mean(corrupt_values))
|
||||
row[f"corrupt_{metric}_sd"] = float(np.std(corrupt_values, ddof=1)) if len(corrupt_values) > 1 else 0.0
|
||||
row["worst_condition_macro_f1"] = condition_f1[worst_condition]
|
||||
row["worst_single_run_macro_f1"] = min(
|
||||
item["macro_f1"] for seed in seeds for item in corrupt_by_seed[seed]
|
||||
)
|
||||
text_30 = [by_seed_condition[(seed, "text_30")] for seed in seeds]
|
||||
row["text_30_macro_f1"] = float(np.mean([item["macro_f1"] for item in text_30]))
|
||||
row["text_30_macro_f1_sd"] = float(np.std([item["macro_f1"] for item in text_30], ddof=1)) if len(text_30) > 1 else 0.0
|
||||
for condition in ("audio_30", "vision_30", "audio_vision_30", "all_modalities_30"):
|
||||
values = [by_seed_condition[(seed, condition)]["macro_f1"] for seed in seeds]
|
||||
row[f"{condition}_macro_f1"] = float(np.mean(values))
|
||||
row[f"{condition}_macro_f1_sd"] = float(np.std(values, ddof=1)) if len(values) > 1 else 0.0
|
||||
row["corrupt_macro_f1"] = row["corrupt_macro_f1_mean"]
|
||||
row["corrupt_mae"] = row["corrupt_mae_mean"]
|
||||
row["corrupt_pearson"] = row["corrupt_pearson_mean"]
|
||||
summaries.append(row)
|
||||
return summaries
|
||||
|
||||
|
||||
def _bootstrap_distributions(
|
||||
method: str,
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]],
|
||||
valid: Split,
|
||||
seeds: list[int],
|
||||
conditions: list[str],
|
||||
group_counts: np.ndarray,
|
||||
) -> dict[str, np.ndarray]:
|
||||
group_names = sorted({sample_id.split("$_$", 1)[0] for sample_id in valid.ids})
|
||||
group_index = {name: index for index, name in enumerate(group_names)}
|
||||
row_group = np.asarray([group_index[sample_id.split("$_$", 1)[0]] for sample_id in valid.ids], dtype=np.int64)
|
||||
n_groups = len(group_names)
|
||||
n_slots = len(seeds) * len(conditions)
|
||||
confusion_by_group = np.zeros((n_groups, n_slots, 9), dtype=np.float64)
|
||||
regression_by_group = np.zeros((n_groups, n_slots, 7), dtype=np.float64)
|
||||
for seed_index, seed in enumerate(seeds):
|
||||
for condition_index, condition in enumerate(conditions):
|
||||
slot = seed_index * len(conditions) + condition_index
|
||||
pred = predictions[(method, seed, condition)]
|
||||
predicted_class = pred["logits"].argmax(axis=-1)
|
||||
code = valid.y_cls * 3 + predicted_class
|
||||
np.add.at(confusion_by_group[:, slot, :], (row_group, code), 1.0)
|
||||
intensity = np.clip(pred["intensity"].reshape(-1), -3.0, 3.0)
|
||||
values = np.stack((
|
||||
np.ones(valid.n),
|
||||
np.abs(valid.y_reg - intensity),
|
||||
valid.y_reg,
|
||||
valid.y_reg ** 2,
|
||||
intensity,
|
||||
intensity ** 2,
|
||||
valid.y_reg * intensity,
|
||||
), axis=-1)
|
||||
for statistic in range(values.shape[-1]):
|
||||
np.add.at(regression_by_group[:, slot, statistic], row_group, values[:, statistic])
|
||||
|
||||
weighted_confusion = np.einsum("rg,gsk->rsk", group_counts, confusion_by_group, optimize=True)
|
||||
cm = weighted_confusion.reshape(len(group_counts), len(seeds), len(conditions), 3, 3)
|
||||
true_count = cm.sum(axis=-1)
|
||||
predicted_count = cm.sum(axis=-2)
|
||||
true_positive = np.diagonal(cm, axis1=-2, axis2=-1)
|
||||
denominator = true_count + predicted_count
|
||||
class_f1 = np.divide(2.0 * true_positive, denominator, out=np.zeros_like(true_positive), where=denominator > 0)
|
||||
macro_f1 = class_f1.mean(axis=-1)
|
||||
|
||||
weighted_regression = np.einsum("rg,gsk->rsk", group_counts, regression_by_group, optimize=True)
|
||||
regression = weighted_regression.reshape(len(group_counts), len(seeds), len(conditions), 7)
|
||||
count = np.maximum(regression[..., 0], 1.0)
|
||||
mae = regression[..., 1] / count
|
||||
sum_y, sum_y2, sum_pred, sum_pred2, sum_yp = (regression[..., index] for index in range(2, 7))
|
||||
covariance = sum_yp - sum_y * sum_pred / count
|
||||
variance_y = np.maximum(sum_y2 - sum_y ** 2 / count, 0.0)
|
||||
variance_pred = np.maximum(sum_pred2 - sum_pred ** 2 / count, 0.0)
|
||||
denominator_corr = np.sqrt(variance_y * variance_pred)
|
||||
pearson = np.divide(covariance, denominator_corr, out=np.zeros_like(covariance), where=denominator_corr > 1e-12)
|
||||
text_30_index = conditions.index("text_30")
|
||||
return {
|
||||
"corrupt_macro_f1": macro_f1[:, :, 1:].mean(axis=(1, 2)),
|
||||
"worst_condition_macro_f1": macro_f1[:, :, 1:].mean(axis=1).min(axis=1),
|
||||
"text_30_macro_f1": macro_f1[:, :, text_30_index].mean(axis=1),
|
||||
"corrupt_mae": mae[:, :, 1:].mean(axis=(1, 2)),
|
||||
"corrupt_pearson": pearson[:, :, 1:].mean(axis=(1, 2)),
|
||||
}
|
||||
|
||||
|
||||
def _paired_bootstrap(
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]],
|
||||
valid: Split,
|
||||
seeds: list[int],
|
||||
conditions: list[str],
|
||||
reps: int,
|
||||
bootstrap_seed: int,
|
||||
summaries: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
groups = sorted({sample_id.split("$_$", 1)[0] for sample_id in valid.ids})
|
||||
rng = np.random.default_rng(bootstrap_seed)
|
||||
draws = rng.integers(0, len(groups), size=(reps, len(groups)))
|
||||
group_counts = np.zeros((reps, len(groups)), dtype=np.float64)
|
||||
for rep in range(reps):
|
||||
group_counts[rep] = np.bincount(draws[rep], minlength=len(groups))
|
||||
candidate = _bootstrap_distributions(MOFE7_MLP, predictions, valid, seeds, conditions, group_counts)
|
||||
reference = _bootstrap_distributions(EARLYCONCAT, predictions, valid, seeds, conditions, group_counts)
|
||||
summary_map = {row["method"]: row for row in summaries}
|
||||
point_keys = {
|
||||
"corrupt_macro_f1": "corrupt_macro_f1",
|
||||
"worst_condition_macro_f1": "worst_condition_macro_f1",
|
||||
"text_30_macro_f1": "text_30_macro_f1",
|
||||
"corrupt_mae": "corrupt_mae",
|
||||
"corrupt_pearson": "corrupt_pearson",
|
||||
}
|
||||
rows = []
|
||||
for metric in SUMMARY_METRICS:
|
||||
delta = candidate[metric] - reference[metric]
|
||||
key = point_keys[metric]
|
||||
rows.append({
|
||||
"comparison": "MoFE-7 MLP vs EarlyConcat",
|
||||
"candidate": MOFE7_MLP,
|
||||
"reference": EARLYCONCAT,
|
||||
"metric": metric,
|
||||
"delta_candidate_minus_reference": float(summary_map[MOFE7_MLP][key] - summary_map[EARLYCONCAT][key]),
|
||||
"bootstrap_ci_2p5": float(np.quantile(delta, 0.025)),
|
||||
"bootstrap_ci_97p5": float(np.quantile(delta, 0.975)),
|
||||
"bootstrap_probability_delta_gt_0": float(np.mean(delta > 0.0)),
|
||||
"bootstrap_replicates": reps,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
"seed": bootstrap_seed,
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def _plot_summary(output: Path, summaries: list[dict[str, Any]]) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
labels = ["EarlyConcat + BiGRU", "MoFE-7 + MLP Router"]
|
||||
by_method = {row["method"]: row for row in summaries}
|
||||
methods = (EARLYCONCAT, MOFE7_MLP)
|
||||
metrics = ("clean_macro_f1", "corrupt_macro_f1", "worst_condition_macro_f1")
|
||||
names = ("Clean", "Mean corrupted", "Worst condition")
|
||||
x = np.arange(len(names))
|
||||
width = 0.34
|
||||
fig, ax = plt.subplots(figsize=(8.6, 4.8), constrained_layout=True)
|
||||
for offset, method, label, color in (
|
||||
(-width / 2, methods[0], labels[0], "#4e79a7"),
|
||||
(width / 2, methods[1], labels[1], "#f28e2b"),
|
||||
):
|
||||
values = [by_method[method][metric] for metric in metrics]
|
||||
ax.bar(x + offset, values, width, label=label, color=color)
|
||||
ax.set_xticks(x, names)
|
||||
ax.set_ylabel("Macro-F1")
|
||||
ax.set_ylim(0, 1)
|
||||
ax.set_title("Q2 selected-model validation comparison")
|
||||
ax.legend(frameon=False)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
fig.savefig(output / "comparison_earlyconcat_mofe7.png", dpi=180)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def _parameter_rows(dims: tuple[int, int, int], device: torch.device) -> list[dict[str, Any]]:
|
||||
models: dict[str, nn.Module] = {
|
||||
EARLYCONCAT: AlignedFusionModel("concat", dims=dims),
|
||||
MOFE7_MLP: MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG),
|
||||
}
|
||||
baseline_count = sum(parameter.numel() for parameter in models[EARLYCONCAT].parameters() if parameter.requires_grad)
|
||||
rows = []
|
||||
for name, model in models.items():
|
||||
count = sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad)
|
||||
rows.append({
|
||||
"method": name,
|
||||
"trainable_parameters": count,
|
||||
"ratio_to_earlyconcat": count / baseline_count,
|
||||
"within_2x_earlyconcat": bool(count <= 2 * baseline_count),
|
||||
})
|
||||
return rows
|
||||
|
||||
|
||||
def _smoke_test(train: Split, output: Path, device: torch.device, seed: int) -> dict[str, Any]:
|
||||
seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
count = min(4, train.n)
|
||||
xs = tuple(torch.as_tensor(x[:count], dtype=torch.float32, device=device) for x in train.x)
|
||||
masks = torch.as_tensor(train.mask[:count].copy(), dtype=torch.bool, device=device)
|
||||
masks[0] = True
|
||||
if count > 1:
|
||||
masks[1, 5:12, 0] = False
|
||||
if count > 2:
|
||||
masks[2, 18:23, :] = False
|
||||
target_class = torch.as_tensor(train.y_cls[:count], dtype=torch.long, device=device)
|
||||
target_intensity = torch.as_tensor(train.y_reg[:count], dtype=torch.float32, device=device)
|
||||
reports: dict[str, Any] = {}
|
||||
models: dict[str, nn.Module] = {
|
||||
EARLYCONCAT: AlignedFusionModel("concat", dims=dims).to(device),
|
||||
MOFE7_MLP: MixtureOfFusionExperts(dims=dims, **MODEL_CONFIG).to(device),
|
||||
}
|
||||
for name, model in models.items():
|
||||
model.train()
|
||||
result = model(xs, masks)
|
||||
loss = _loss(result, target_class, target_intensity)
|
||||
loss.backward()
|
||||
gradient = sum(float(p.grad.detach().abs().sum().cpu()) for p in model.parameters() if p.grad is not None)
|
||||
reports[name] = {
|
||||
"logits_shape": list(result["logits"].shape),
|
||||
"intensity_shape": list(result["intensity"].shape),
|
||||
"finite_loss": bool(torch.isfinite(loss).item()),
|
||||
"gradient_l1": gradient,
|
||||
}
|
||||
mofe_model = models[MOFE7_MLP]
|
||||
mo = mofe_model(xs, masks)
|
||||
active = mo["availability"].any(dim=-1)
|
||||
alpha_sums = mo["alpha"].sum(dim=-1)
|
||||
alpha_error = float((alpha_sums[active] - 1).abs().max().cpu()) if active.any() else 0.0
|
||||
unavailable_weights = float(mo["alpha"].masked_select(~mo["availability"]).abs().max().cpu()) if (~mo["availability"]).any() else 0.0
|
||||
expert_gradients = {
|
||||
name: sum(float(parameter.grad.detach().abs().sum().cpu()) for parameter in expert.parameters() if parameter.grad is not None)
|
||||
for name, expert in mofe_model.experts.items()
|
||||
}
|
||||
router_gradient = sum(float(parameter.grad.detach().abs().sum().cpu()) for parameter in mofe_model.router.parameters() if parameter.grad is not None)
|
||||
if alpha_error > 1e-6 or unavailable_weights > 1e-8:
|
||||
raise RuntimeError(f"MoFE routing mask invariant failed: sum_error={alpha_error}, unavailable={unavailable_weights}")
|
||||
if not all(value > 0 for value in expert_gradients.values()) or router_gradient <= 0:
|
||||
raise RuntimeError(f"MoFE expert/router gradients are incomplete: {expert_gradients}; router={router_gradient}")
|
||||
report = {
|
||||
"passed": all(item["finite_loss"] and item["gradient_l1"] > 0 for item in reports.values()),
|
||||
"seed": seed,
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"batch_size_checked": count,
|
||||
"steps": train.steps,
|
||||
"models": reports,
|
||||
"mofe_experts": list(EXPERT_NAMES),
|
||||
"mofe_alpha_shape": list(mo["alpha"].shape),
|
||||
"mofe_max_weight_sum_error": alpha_error,
|
||||
"mofe_max_weight_on_unavailable_experts": unavailable_weights,
|
||||
"mofe_expert_gradient_l1": expert_gradients,
|
||||
"mofe_router_gradient_l1": router_gradient,
|
||||
"parameter_count": {row["method"]: row["trainable_parameters"] for row in _parameter_rows(dims, device)},
|
||||
}
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
(output / "smoke_test.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
|
||||
return report
|
||||
|
||||
|
||||
def _run(args: argparse.Namespace) -> None:
|
||||
output = args.output_dir.resolve()
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
device = _device_for(args.device)
|
||||
torch.set_num_threads(args.threads)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
raw = load_aligned()
|
||||
computed_stats = fit_robust_stats(raw["train"])
|
||||
reference_stats_path = REFERENCE_OUTPUT / "aligned_robust_stats.npz"
|
||||
if reference_stats_path.exists():
|
||||
stats = RobustStats.load(reference_stats_path)
|
||||
scaler_diff = max(
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed_stats.center, stats.center)),
|
||||
max(float(np.max(np.abs(a - b))) for a, b in zip(computed_stats.scale, stats.scale)),
|
||||
)
|
||||
else:
|
||||
stats = computed_stats
|
||||
scaler_diff = 0.0
|
||||
train = apply_robust_stats(raw["train"], stats)
|
||||
valid = apply_robust_stats(raw["valid"], stats)
|
||||
stats.save(output / "aligned_robust_stats.npz")
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
feature_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
if not feature_path.exists():
|
||||
raise FileNotFoundError(f"official aligned feature file not found: {feature_path}")
|
||||
|
||||
if args.phase == "smoke":
|
||||
report = _smoke_test(train, output, device, args.seeds[0])
|
||||
report["scaler_max_abs_difference_from_reference"] = scaler_diff
|
||||
(output / "smoke_test.json").write_text(json.dumps(report, indent=2), encoding="utf-8")
|
||||
print(f"selected-model smoke: passed={report['passed']} device={device}", flush=True)
|
||||
return
|
||||
|
||||
seeds = list(args.seeds)
|
||||
metrics_rows: list[dict[str, Any]] = []
|
||||
predictions: dict[tuple[str, int, str], dict[str, np.ndarray]] = {}
|
||||
router_rows: list[dict[str, Any]] = []
|
||||
utility_rows: list[dict[str, Any]] = []
|
||||
expert_rows: list[dict[str, Any]] = []
|
||||
history_rows: list[dict[str, Any]] = []
|
||||
best_epochs: dict[str, int] = {}
|
||||
condition_names: list[str] = []
|
||||
|
||||
for seed in seeds:
|
||||
baseline_dir = output / "models" / "baselines" / "concat" / f"seed_{seed}"
|
||||
baseline, baseline_epoch, baseline_history = _load_or_train_concat(
|
||||
train, valid, baseline_dir, device, seed, args.epochs, args.patience,
|
||||
args.batch_size, args.reuse_checkpoints and not args.force_retrain,
|
||||
)
|
||||
mofe_dir = output / "models" / MOFE7_MLP / f"seed_{seed}"
|
||||
mofe, mofe_epoch, mofe_history = _train_mofe(
|
||||
train, valid, mofe_dir, device, seed, args.epochs, args.patience,
|
||||
args.batch_size, args.reuse_checkpoints and not args.force_retrain,
|
||||
)
|
||||
best_epochs[f"{EARLYCONCAT}_seed_{seed}"] = baseline_epoch
|
||||
best_epochs[f"{MOFE7_MLP}_seed_{seed}"] = mofe_epoch
|
||||
history_rows.extend(baseline_history)
|
||||
history_rows.extend(mofe_history)
|
||||
|
||||
conditions = _conditions(valid, seed)
|
||||
names = [condition for condition, _, _ in conditions]
|
||||
if condition_names and names != condition_names:
|
||||
raise RuntimeError("validation condition ordering changed between seeds")
|
||||
condition_names = names
|
||||
for method, model in ((EARLYCONCAT, baseline), (MOFE7_MLP, mofe)):
|
||||
for condition, rate, masks in conditions:
|
||||
prediction = _predict(model, valid, masks, device, args.batch_size)
|
||||
predictions[(method, seed, condition)] = prediction
|
||||
metrics_rows.append(_condition_row(method, seed, condition, rate, valid, prediction))
|
||||
if method == MOFE7_MLP:
|
||||
route_row, utility_row = _diagnostics(seed, condition, masks, prediction)
|
||||
router_rows.append(route_row)
|
||||
utility_rows.append(utility_row)
|
||||
for expert in EXPERT_NAMES:
|
||||
forced = _predict(model, valid, masks, device, args.batch_size, force_expert=expert)
|
||||
observed = masks[..., list(SUBSETS[expert])].all(axis=-1)
|
||||
expert_rows.append({
|
||||
"method": MOFE7_MLP,
|
||||
"seed": seed,
|
||||
"condition": condition,
|
||||
"expert": expert,
|
||||
"available_position_fraction": float(observed.mean()),
|
||||
**_metric_dict(valid.y_cls, valid.y_reg, forced["logits"], forced["intensity"]),
|
||||
})
|
||||
print(f"evaluated {method}/seed{seed}", flush=True)
|
||||
del baseline, mofe
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summaries = _summary_rows(metrics_rows)
|
||||
paired = _paired_bootstrap(
|
||||
predictions, valid, seeds, condition_names, args.bootstrap_reps,
|
||||
args.bootstrap_seed, summaries,
|
||||
) if args.bootstrap_reps > 0 else []
|
||||
parameter_rows = _parameter_rows(dims, device)
|
||||
output_rows = {
|
||||
"metrics_by_condition.csv": metrics_rows,
|
||||
"summary.csv": summaries,
|
||||
"paired_bootstrap.csv": paired,
|
||||
"parameter_count.csv": parameter_rows,
|
||||
"router_weights_by_condition.csv": router_rows,
|
||||
"routing_entropy.csv": router_rows,
|
||||
"modality_utility_by_condition.csv": utility_rows,
|
||||
"expert_condition_matrix.csv": expert_rows,
|
||||
"training_history.csv": history_rows,
|
||||
}
|
||||
for filename, rows in output_rows.items():
|
||||
_write_csv(output / filename, rows)
|
||||
_plot_summary(output / "figures", summaries)
|
||||
|
||||
manifest = {
|
||||
"experiment": "Q2 selected models: EarlyConcat + BiGRU and MoFE-7 + MLP Router",
|
||||
"created_unix": time.time(),
|
||||
"python_version": sys.version,
|
||||
"torch_version": torch.__version__,
|
||||
"numpy_version": np.__version__,
|
||||
"device": str(device),
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"feature_file": str(feature_path),
|
||||
"feature_sha256": _sha256(feature_path),
|
||||
"feature_dimensions": dict(zip(MODALITIES, dims)),
|
||||
"sequence_length": train.steps,
|
||||
"representation_note": "official ordered 50-wordpiece positions; not 50 physical-time bins",
|
||||
"train_examples": train.n,
|
||||
"valid_examples": valid.n,
|
||||
"train_source_video_groups": len({sample_id.split("$_$", 1)[0] for sample_id in train.ids}),
|
||||
"valid_source_video_groups": len({sample_id.split("$_$", 1)[0] for sample_id in valid.ids}),
|
||||
"train_only_scaler": str(output / "aligned_robust_stats.npz"),
|
||||
"scaler_max_abs_difference_from_reference": scaler_diff,
|
||||
"test_labels_used": False,
|
||||
"seeds": seeds,
|
||||
"epochs_max": args.epochs,
|
||||
"patience": args.patience,
|
||||
"batch_size": args.batch_size,
|
||||
"optimizer": "AdamW(lr=1.5e-4, weight_decay=1e-4), gradient clip 1.0",
|
||||
"training_mask_augmentation": "same contiguous-block augment_masks protocol for both models",
|
||||
"validation_conditions": condition_names,
|
||||
"validation_corruption_seed": "seed + 13 + pattern_index*101 + int(rate*1000)",
|
||||
"loss": "cross_entropy + 0.5*SmoothL1(intensity/3, regression_label/3)",
|
||||
"models": {
|
||||
EARLYCONCAT: "project modalities independently, concatenate features and masks, then BiGRU",
|
||||
MOFE7_MLP: {
|
||||
"experts": list(EXPERT_NAMES),
|
||||
"router": "MLP over per-position observed values and local observation statistics",
|
||||
"availability": "hard mask; unavailable expert weights are zero",
|
||||
"shared_temporal_backbone": "one BiGRU after position-wise expert mixture",
|
||||
},
|
||||
},
|
||||
"best_epochs": best_epochs,
|
||||
"paired_bootstrap": {
|
||||
"replicates": args.bootstrap_reps,
|
||||
"seed": args.bootstrap_seed,
|
||||
"resampling_unit": "source video id",
|
||||
"paired": True,
|
||||
},
|
||||
}
|
||||
(output / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
|
||||
print(f"selected-model results saved to {output}", flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Train and compare the two retained Q2 models.")
|
||||
parser.add_argument("--phase", choices=("smoke", "full"), default="full")
|
||||
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("--seeds", type=int, nargs="+", default=list(SEEDS))
|
||||
parser.add_argument("--bootstrap-reps", type=int, default=1000)
|
||||
parser.add_argument("--bootstrap-seed", type=int, default=20260924)
|
||||
parser.add_argument("--reuse-checkpoints", action="store_true")
|
||||
parser.add_argument("--force-retrain", action="store_true")
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
|
||||
_run(parser.parse_args())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user