Files
modeling_zhaocui/deep_learning/Q2/q2/train_mofe.py
T

788 lines
35 KiB
Python

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()