"""Train Q3 on the official training split and explain Attachment 4 cases.""" from __future__ import annotations import argparse import csv import hashlib import json import math import random import time from pathlib import Path from typing import Any import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import torch import torch.nn.functional as F from sklearn.metrics import accuracy_score, confusion_matrix, f1_score, mean_absolute_error, mean_squared_error from transformers import AutoTokenizer from ..adapter import Q1AlignmentAdapter from ..data_paths import ATTACHMENT4, DATA_ROOT, PROJECT_ROOT from ..model.early_concat import AlignedFusionModel from ..q2.deep_learning.q2.evaluate_math_protocol import continuous_mask, scenario_seed from ..q2.math.data import ( MODALITIES, fit_preprocessor, load_official_splits, restricted_load, transform_split, ) SEED = 20260924 TEXT_MODEL_ID = "google-bert/bert-base-uncased" CLASS_NAMES = ("negative", "neutral", "positive") MODALITY_NAMES = ("text", "audio", "vision") def _sha256(path: Path) -> str: h = hashlib.sha256() with path.open("rb") as stream: for block in iter(lambda: stream.read(1024 * 1024), b""): h.update(block) return h.hexdigest() def _decode(value: Any) -> str: if isinstance(value, bytes): return value.decode("utf-8", errors="replace") if isinstance(value, np.bytes_): return bytes(value).decode("utf-8", errors="replace") if isinstance(value, np.ndarray): if value.shape == (): return _decode(value.item()) return " ".join(_decode(x) for x in value.reshape(-1)) return str(value) def _scalar_int(value: Any, field: str) -> int: arr = np.asarray(value).reshape(-1) if not len(arr): raise ValueError(f"Attachment 4 {field} is empty") return int(arr[0]) def _attachment4_location(version: str) -> tuple[Path, Path]: inner = ATTACHMENT4 / "附件4-可解释专项视频样本与特征文件" version_dir = inner / ("未对齐版本" if version == "unaligned_50" else "对齐版本") video_dir = inner / "videos" if not version_dir.is_dir(): raise FileNotFoundError(f"Attachment 4 {version} directory not found: {version_dir}") return version_dir, video_dir def _read_attachment4(version: str) -> tuple[list[dict[str, Any]], dict[str, str]]: if version != "unaligned_50": raise ValueError("Q3 explanation currently uses the official unaligned_50 Attachment 4 features") version_dir, video_dir = _attachment4_location(version) paths = sorted(version_dir.glob("*.pkl"), key=lambda p: p.name) if len(paths) != 20: raise FileNotFoundError(f"expected 20 Attachment 4 cases, found {len(paths)} under {version_dir}") video_by_stem = {p.stem: p for p in video_dir.rglob("*.mp4")} if video_dir.is_dir() else {} adapter = Q1AlignmentAdapter(target_steps=50) cases: list[dict[str, Any]] = [] for path in paths: raw = restricted_load(path) case_id = _decode(raw.get("id", path.stem)).strip() or path.stem text_bert = np.asarray(raw["text_bert"], dtype=np.int64) if text_bert.ndim == 3 and text_bert.shape[0] == 1: text_bert = text_bert[0] if text_bert.shape != (3, 50): raise ValueError(f"{path.name}: expected text_bert (3,50), got {text_bert.shape}") record = { "id": case_id, "sequence_order_verified": True, "attention_mask": text_bert[1].astype(bool), "text": np.asarray(raw["text"], dtype=np.float32), "audio": np.asarray(raw["audio"], dtype=np.float32), "vision": np.asarray(raw["vision"], dtype=np.float32), "audio_length": _scalar_int(raw["audio_lengths"], "audio_lengths"), "vision_length": _scalar_int(raw["vision_lengths"], "vision_lengths"), } aligned = adapter.align(record, mode="relative") mask = np.stack([aligned.observed[m] for m in MODALITIES], axis=-1) features = {m: aligned.features[m].astype(np.float32) for m in MODALITIES} transcript = _decode(raw.get("raw_text", "")) video_path = video_by_stem.get(path.stem) or video_by_stem.get(case_id) media = "" if video_path is not None: try: media = video_path.resolve().relative_to(DATA_ROOT).as_posix() except ValueError: media = str(video_path.resolve()) cases.append({ "case_id": case_id, "source_file": path, "source_sha256": _sha256(path), "transcript": transcript, "text_bert": text_bert, "raw": raw, "features": features, "mask": mask, "target_intervals": aligned.target_intervals.astype(np.float32), "provenance": aligned.provenance, "video_path": media, "coordinate_mode": aligned.metadata["coordinate_mode"], "input_audit": { "case_id": case_id, "source_file": path.name, "source_sha256": _sha256(path), "coordinate_mode": aligned.metadata["coordinate_mode"], "physical_time_alignment": False, "audio_reported_length": record["audio_length"], "vision_reported_length": record["vision_length"], "audio_length_conflict": bool(aligned.provenance["audio"].length_conflict), "vision_length_conflict": bool(aligned.provenance["vision"].length_conflict), "text_visible_target_slots": int(aligned.observed["text"].sum()), "audio_visible_target_slots": int(aligned.observed["audio"].sum()), "vision_visible_target_slots": int(aligned.observed["vision"].sum()), "source_video": media, }, }) return cases, {"version_dir": str(version_dir), "video_dir": str(video_dir)} def _split_arrays(split: Any, transformed: dict[str, np.ndarray]) -> tuple[tuple[np.ndarray, ...], np.ndarray]: return tuple(transformed[m] for m in MODALITIES), np.asarray(split.mask, dtype=bool) def _predict( model: torch.nn.Module, xs: tuple[np.ndarray, ...], masks: np.ndarray, device: torch.device, batch_size: int, ) -> dict[str, np.ndarray]: model.eval() logits: list[np.ndarray] = [] intensity: list[np.ndarray] = [] with torch.inference_mode(): for start in range(0, len(masks), batch_size): end = min(start + batch_size, len(masks)) batch_x = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs) batch_mask = torch.as_tensor(masks[start:end], dtype=torch.bool, device=device) output = model(batch_x, batch_mask) logits.append(output["logits"].float().cpu().numpy()) intensity.append(output["intensity"].float().cpu().numpy()) return {"logits": np.concatenate(logits), "intensity": np.concatenate(intensity)} def _metrics(y_cls: np.ndarray, y_reg: np.ndarray, prediction: dict[str, np.ndarray]) -> dict[str, Any]: logits = np.asarray(prediction["logits"]) score = np.clip(np.asarray(prediction["intensity"]).reshape(-1), -3.0, 3.0) predicted = logits.argmax(axis=-1) pearson = float(np.corrcoef(y_reg, score)[0, 1]) if np.std(y_reg) > 0 and np.std(score) > 0 else None return { "n": int(len(y_cls)), "accuracy": float(accuracy_score(y_cls, predicted)), "macro_f1": float(f1_score(y_cls, predicted, labels=[0, 1, 2], average="macro", zero_division=0)), "mae": float(mean_absolute_error(y_reg, score)), "rmse": float(math.sqrt(mean_squared_error(y_reg, score))), "pearson": pearson, "confusion_matrix_rows_true_columns_predicted": confusion_matrix(y_cls, predicted, labels=[0, 1, 2]).tolist(), "per_class_support": {CLASS_NAMES[i]: int(np.sum(y_cls == i)) for i in range(3)}, } def _loss(logits: torch.Tensor, intensity: torch.Tensor, y_cls: torch.Tensor, y_reg: torch.Tensor) -> torch.Tensor: return F.cross_entropy(logits, y_cls) + 0.5 * F.smooth_l1_loss(intensity / 3.0, y_reg / 3.0) def _validation_loss( model: torch.nn.Module, xs: tuple[np.ndarray, ...], masks: list[np.ndarray], y_cls: np.ndarray, y_reg: np.ndarray, device: torch.device, batch_size: int, ) -> float: values: list[float] = [] model.eval() with torch.inference_mode(): for scenario in masks: total, count = 0.0, 0 for start in range(0, len(y_cls), batch_size): end = min(start + batch_size, len(y_cls)) bx = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs) bm = torch.as_tensor(scenario[start:end], dtype=torch.bool, device=device) by = torch.as_tensor(y_cls[start:end], dtype=torch.long, device=device) br = torch.as_tensor(y_reg[start:end], dtype=torch.float32, device=device) out = model(bx, bm) total += float(_loss(out["logits"], out["intensity"], by, br).item()) * (end - start) count += end - start values.append(total / max(1, count)) return float(np.mean(values)) def _train(args: argparse.Namespace, out_dir: Path) -> tuple[AlignedFusionModel, dict[str, Any], dict[str, Any]]: feature_path: Path if args.data_path is not None: feature_path = args.data_path.expanduser().resolve() else: from ..data_paths import ATTACHMENT2 feature_path = ATTACHMENT2 / f"{args.input_version}.pkl" if not feature_path.is_file(): raise FileNotFoundError(f"Q3 training feature file not found: {feature_path}") raw_splits = load_official_splits(feature_path, version=args.input_version) train = raw_splits["train"] valid = raw_splits["valid"] fitted = fit_preprocessor(train) transformed = {name: transform_split(split, fitted) for name, split in raw_splits.items()} train_x, train_mask = _split_arrays(train, transformed["train"]) valid_x, valid_mask = _split_arrays(valid, transformed["valid"]) dims = tuple(int(x.shape[-1]) for x in train_x) np.savez_compressed(out_dir / "preprocessor.npz", **{ f"{modality}_{key}": value for modality, state in fitted.items() for key, value in state.items() }) device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available() else ("cpu" if args.device == "auto" else args.device)) if device.type == "cuda": torch.cuda.manual_seed_all(SEED) random.seed(SEED) np.random.seed(SEED) torch.manual_seed(SEED) torch.set_num_threads(4) model = AlignedFusionModel("concat", dims=dims).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay) y_cls = np.asarray(train.class_y, dtype=np.int64) y_reg = np.asarray(train.regression_y, dtype=np.float32) vy_cls = np.asarray(valid.class_y, dtype=np.int64) vy_reg = np.asarray(valid.regression_y, dtype=np.float32) valid_rng_masks: list[np.ndarray] = [valid_mask.copy()] for rate, mode in ((0.3, "single"), (0.3, "sync"), (0.5, "async")): key = f"{rate:.1f}/{mode}" valid_rng_masks.append(np.stack([ continuous_mask(mask, rate, mode, np.random.default_rng(scenario_seed(SEED + 177, sid, key))) for sid, mask in zip(valid.ids, valid_mask) ])) best = float("inf") best_epoch = 0 stale = 0 history: list[dict[str, Any]] = [] for epoch in range(1, args.epochs + 1): model.train() train_corruption = np.stack([ continuous_mask( mask, float(np.random.choice((0.0, 0.1, 0.3, 0.5, 0.7))), str(np.random.choice(("single", "sync", "partial", "async"))), np.random.default_rng(scenario_seed(SEED + epoch, sid, f"train/{epoch}")), ) for sid, mask in zip(train.ids, train_mask) ]) order = np.random.permutation(len(y_cls)) losses: list[float] = [] for start in range(0, len(order), args.batch_size): ix = order[start:start + args.batch_size] bx = tuple(torch.as_tensor(x[ix], dtype=torch.float32, device=device) for x in train_x) bm = torch.as_tensor(train_corruption[ix], dtype=torch.bool, device=device) by = torch.as_tensor(y_cls[ix], dtype=torch.long, device=device) br = torch.as_tensor(y_reg[ix], dtype=torch.float32, device=device) optimizer.zero_grad(set_to_none=True) output = model(bx, bm) loss = _loss(output["logits"], output["intensity"], by, br) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() losses.append(float(loss.item())) validation = _validation_loss(model, valid_x, valid_rng_masks, vy_cls, vy_reg, device, args.batch_size) history.append({"epoch": epoch, "train_loss": float(np.mean(losses)), "selection_loss": validation}) print(f"Q3 epoch {epoch}/{args.epochs}: train={np.mean(losses):.5f}, validation={validation:.5f}", flush=True) if validation < best - 1e-7: best, best_epoch, stale = validation, epoch, 0 torch.save({"state_dict": model.state_dict(), "dims": dims, "seed": SEED, "best_epoch": epoch}, out_dir / "model_best.pt") else: stale += 1 if stale >= args.patience: break checkpoint = torch.load(out_dir / "model_best.pt", map_location=device, weights_only=True) model.load_state_dict(checkpoint["state_dict"]) model.eval() prediction = _predict(model, valid_x, valid_mask, device, args.batch_size) metric = _metrics(vy_cls, vy_reg, prediction) metric["best_epoch"] = best_epoch metric["selection_loss_clean_plus_fixed_missing_scenarios"] = best metric["input_version"] = args.input_version metric["adapter"] = "Q1AlignmentAdapter relative normalized progress" metric["physical_time_alignment"] = False _write_csv(out_dir / "training_history.csv", history) _write_json(out_dir / "validation_metrics.json", metric) validation_rows = [] for i, sid in enumerate(valid.ids): prob = torch.softmax(torch.as_tensor(prediction["logits"][i]), dim=-1).numpy() validation_rows.append({ "sample_id": sid, "true_class": int(vy_cls[i]), "true_class_name": CLASS_NAMES[int(vy_cls[i])], "true_sentiment": float(vy_reg[i]), "predicted_class": int(prob.argmax()), "predicted_class_name": CLASS_NAMES[int(prob.argmax())], "predicted_sentiment": float(prediction["intensity"][i]), "p_negative": float(prob[0]), "p_neutral": float(prob[1]), "p_positive": float(prob[2]), "absolute_error": float(abs(vy_reg[i] - prediction["intensity"][i])), }) _write_csv(out_dir / "validation_predictions.csv", validation_rows) errors = sorted( (row for row in validation_rows if row["true_class"] != row["predicted_class"] or row["absolute_error"] >= metric["mae"]), key=lambda row: (-row["absolute_error"], row["sample_id"]), ) _write_csv(out_dir / "validation_errors.csv", errors[:100]) return model, {"metrics": metric, "feature_sha256": _sha256(feature_path), "feature_path": str(feature_path)}, {"x": valid_x, "mask": valid_mask, "y_cls": vy_cls, "y_reg": vy_reg, "prediction": prediction} def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None: if not rows: return path.parent.mkdir(parents=True, exist_ok=True) fields = list(dict.fromkeys(key for row in rows for key in row)) with path.open("w", encoding="utf-8-sig", newline="") as stream: writer = csv.DictWriter(stream, fieldnames=fields) writer.writeheader() writer.writerows(rows) def _write_json(path: Path, payload: Any) -> None: path.write_text(json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8") def _model_output(model: torch.nn.Module, xs: tuple[torch.Tensor, ...], mask: torch.Tensor) -> dict[str, torch.Tensor]: model.eval() with torch.inference_mode(): return model(xs, mask) def _span_evidence(case: dict[str, Any], modality_index: int, slot: int, tokenizer: Any) -> dict[str, Any]: modality = MODALITY_NAMES[modality_index] weights = case["provenance"][modality].source_weights.getrow(slot) source_rows = weights.indices.tolist() if source_rows: low, high = min(source_rows), max(source_rows) + 1 else: low = high = 0 start, end = case["target_intervals"][slot].astype(float).tolist() text = "" if modality == "text" and source_rows: ids = np.asarray(case["text_bert"][0], dtype=np.int64) token_ids = [int(ids[i]) for i in source_rows if i < len(ids) and int(ids[i]) not in tokenizer.all_special_ids] text = " ".join(tokenizer.convert_ids_to_tokens(token_ids)) elif modality == "audio": text = f"audio feature rows {low}–{high - 1}; inspect the same relative span in the linked source video/audio" else: text = f"video feature rows {low}–{high - 1}; inspect the same relative span in the linked source video" return { "modality": modality, "slot": int(slot), "relative_start": float(start), "relative_end": float(end), "source_row_start": int(low), "source_row_end_exclusive": int(high), "evidence": text, } def _explain_case( model: torch.nn.Module, case: dict[str, Any], stats: dict[str, dict[str, np.ndarray]], tokenizer: Any, device: torch.device, batch_size: int, ) -> tuple[dict[str, Any], list[dict[str, Any]]]: values: dict[str, np.ndarray] = {} for modality_index, modality in enumerate(MODALITIES): arr = case["features"][modality].astype(np.float32) arr = np.clip((arr - stats[modality]["mean"]) / stats[modality]["std"], -10.0, 10.0) arr[~case["mask"][:, modality_index]] = 0.0 values[modality] = arr xs = tuple(torch.as_tensor(values[m][None], dtype=torch.float32, device=device) for m in MODALITIES) mask = torch.as_tensor(case["mask"][None], dtype=torch.bool, device=device) full = _model_output(model, xs, mask) probs = torch.softmax(full["logits"], dim=-1)[0].cpu().numpy() pred = int(np.argmax(probs)) contributions: dict[str, float] = {} local_rows: list[dict[str, Any]] = [] for m, modality in enumerate(MODALITY_NAMES): ablated_mask = mask.clone() ablated_mask[:, :, m] = False ablated = _model_output(model, xs, ablated_mask) ablated_p = torch.softmax(ablated["logits"], dim=-1)[0, pred].item() contributions[modality] = float(probs[pred] - ablated_p) observed_slots = np.flatnonzero(case["mask"][:, m]) if not len(observed_slots): continue impacts: list[tuple[int, float]] = [] for start in range(0, len(observed_slots), batch_size): chosen = observed_slots[start:start + batch_size] bx = tuple(x.repeat(len(chosen), 1, 1) for x in xs) bm = mask.repeat(len(chosen), 1, 1) row_idx = torch.arange(len(chosen), device=device) slot_idx = torch.as_tensor(chosen, dtype=torch.long, device=device) bm[row_idx, slot_idx, m] = False output = _model_output(model, bx, bm) hidden_p = torch.softmax(output["logits"], dim=-1)[:, pred].cpu().numpy() impacts.extend((int(slot), float(probs[pred] - p)) for slot, p in zip(chosen, hidden_p)) for slot, impact in sorted(impacts, key=lambda row: (-row[1], row[0]))[:3]: evidence = _span_evidence(case, m, slot, tokenizer) evidence["probability_drop"] = impact evidence["case_id"] = case["case_id"] evidence["source_video"] = case["video_path"] local_rows.append(evidence) principal = max(contributions, key=contributions.get) intensity = float(full["intensity"][0].cpu().item()) explanation = { "case_id": case["case_id"], "predicted_class": pred, "predicted_class_name": CLASS_NAMES[pred], "predicted_sentiment": intensity, "p_negative": float(probs[0]), "p_neutral": float(probs[1]), "p_positive": float(probs[2]), "principal_modality": principal, "text_contribution": contributions["text"], "audio_contribution": contributions["audio"], "vision_contribution": contributions["vision"], "transcript": case["transcript"], "source_video": case["video_path"], "coordinate_mode": case["coordinate_mode"], "interpretation_method": "single-modality and single-slot occlusion; probability drops measure model sensitivity", } return explanation, local_rows def _write_cards(out_dir: Path, case_by_id: dict[str, dict[str, Any]], explanations: list[dict[str, Any]], local_rows: list[dict[str, Any]]) -> str: cards = out_dir / "explanation_cards" cards.mkdir(parents=True, exist_ok=True) rows_by_id: dict[str, list[dict[str, Any]]] = {} for row in local_rows: rows_by_id.setdefault(str(row["case_id"]), []).append(row) for item in explanations: evidence = rows_by_id.get(str(item["case_id"]), []) lines = [f"# Q3 Explanation: {item['case_id']}", "", f"- Prediction: **{item['predicted_class_name']}**", f"- Sentiment score: {item['predicted_sentiment']:.3f}", f"- Probabilities (negative / neutral / positive): {item['p_negative']:.3f} / {item['p_neutral']:.3f} / {item['p_positive']:.3f}", f"- Main modality by occlusion: **{item['principal_modality']}**", f"- Source video/audio: `{item['source_video'] or 'not found in the supplied video folder'}`", f"- Coordinate: normalized progress `[0,1]`; no physical timestamps are inferred from the unaligned feature rows.", "", "## Modality contribution", "", "Removing one modality changes the predicted-class probability by the values below. Positive values mean that modality supports the prediction under this model.", "", "| Modality | Probability drop |", "|---|---:|"] for modality in MODALITY_NAMES: lines.append(f"| {modality} | {item[f'{modality}_contribution']:.4f} |") lines.extend(["", "## Local evidence", "", "Local values are single-slot occlusion sensitivity. Audio/video spans are relative positions in the supplied source clip; text is shown as BERT tokens and the full transcript is retained below.", ""]) for evidence_row in evidence: lines.append(f"- **{evidence_row['modality']}**, slots {evidence_row['slot']} `[0-based]`, relative {evidence_row['relative_start']:.3f}–{evidence_row['relative_end']:.3f}, probability drop {evidence_row['probability_drop']:.4f}: {evidence_row['evidence']}") lines.extend(["", "## Transcript", "", item["transcript"] or "(not supplied)", "", "## Interpretation note", "", "Occlusion scores describe how this trained model responds to removing features. They are not causal effects or proof that the signal expresses the named emotion.", ""]) safe = "".join(c if c.isalnum() or c in "-_" else "_" for c in str(item["case_id"])) (cards / f"{safe}.md").write_text("\n".join(lines), encoding="utf-8") confidence = np.asarray([max(row["p_negative"], row["p_neutral"], row["p_positive"]) for row in explanations]) representative = explanations[int(np.argmin(np.abs(confidence - np.median(confidence))))] source = cards / ("".join(c if c.isalnum() or c in "-_" else "_" for c in str(representative["case_id"])) + ".md") representative_card = out_dir / "typical_explanation_card.md" representative_card.write_text(source.read_text(encoding="utf-8"), encoding="utf-8") return str(representative["case_id"]) def _plot_validation(out_dir: Path, y_cls: np.ndarray, prediction: dict[str, np.ndarray]) -> None: pred_cls = prediction["logits"].argmax(axis=-1) matrix = confusion_matrix(y_cls, pred_cls, labels=[0, 1, 2]) fig, axes = plt.subplots(1, 2, figsize=(10, 4), constrained_layout=True) image = axes[0].imshow(matrix, cmap="Blues") axes[0].set_xticks(range(3), CLASS_NAMES, rotation=15) axes[0].set_yticks(range(3), CLASS_NAMES) axes[0].set_xlabel("Predicted") axes[0].set_ylabel("True") axes[0].set_title("Validation confusion matrix") for (i, j), value in np.ndenumerate(matrix): axes[0].text(j, i, str(value), ha="center", va="center") fig.colorbar(image, ax=axes[0], fraction=0.046) axes[1].scatter(prediction["intensity"], prediction["true_sentiment"], s=12, alpha=0.55) axes[1].plot([-3, 3], [-3, 3], color="gray", linestyle="--", linewidth=1) axes[1].set(xlim=(-3, 3), ylim=(-3, 3), xlabel="Predicted sentiment", ylabel="True sentiment", title="Validation intensity") fig.savefig(out_dir / "validation_diagnostics.png", dpi=180) plt.close(fig) def run(args: argparse.Namespace) -> None: out_dir = args.output_dir.expanduser().resolve() out_dir.mkdir(parents=True, exist_ok=True) started = time.time() model, training_info, validation = _train(args, out_dir) _plot_validation(out_dir, validation["y_cls"], {**validation["prediction"], "true_sentiment": validation["y_reg"]}) tokenizer = AutoTokenizer.from_pretrained(TEXT_MODEL_ID, use_fast=True) with np.load(out_dir / "preprocessor.npz", allow_pickle=False) as saved: stats = {m: {key: saved[f"{m}_{key}"].astype(np.float32) for key in ("mean", "std")} for m in MODALITIES} cases, input_locations = _read_attachment4(args.attachment4_version) device = next(model.parameters()).device explanation_rows, all_local = [], [] prediction_rows = [] for case in cases: explanation, local = _explain_case(model, case, stats, tokenizer, device, args.explanation_batch_size) explanation_rows.append(explanation) all_local.extend(local) prediction_rows.append({key: explanation[key] for key in ( "case_id", "predicted_class", "predicted_class_name", "predicted_sentiment", "p_negative", "p_neutral", "p_positive", "source_video", )}) _write_csv(out_dir / "attachment4_predictions.csv", prediction_rows) _write_csv(out_dir / "attachment4_explanations.csv", explanation_rows) _write_csv(out_dir / "attachment4_local_evidence.csv", all_local) _write_csv(out_dir / "attachment4_input_audit.csv", [case["input_audit"] for case in cases]) typical_id = _write_cards(out_dir, {case["case_id"]: case for case in cases}, explanation_rows, all_local) manifest = { "created_at_unix": time.time(), "elapsed_seconds": time.time() - started, "seed": SEED, "training_input": training_info, "attachment4": input_locations, "attachment4_version": args.attachment4_version, "attachment4_cases": len(cases), "adapter": "Q1AlignmentAdapter shared relative-progress projection", "coordinate_limit": "source-time stamps are absent; local audio/video positions are normalized progress, not seconds", "model": "EarlyConcat + BiGRU", "explanation": "single-modality and single-slot occlusion probability drops; model sensitivity, not causal attribution", "validation_metrics": validation["metrics"], "typical_explanation_case": typical_id, "outputs": [ "model_best.pt", "preprocessor.npz", "validation_metrics.json", "validation_predictions.csv", "validation_errors.csv", "validation_diagnostics.png", "attachment4_predictions.csv", "attachment4_explanations.csv", "attachment4_local_evidence.csv", "attachment4_input_audit.csv", "typical_explanation_card.md", ], } _write_json(out_dir / "run_manifest.json", manifest) print(f"Q3 complete: {len(cases)} Attachment 4 predictions saved under {out_dir}", flush=True) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="unaligned_50") parser.add_argument("--attachment4-version", choices=("unaligned_50",), default="unaligned_50") parser.add_argument("--data-path", type=Path, default=None, help="Optional explicit Attachment 2 pickle path") parser.add_argument("--output-dir", type=Path, default=PROJECT_ROOT / "output" / "q3") parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") parser.add_argument("--epochs", type=int, default=12) parser.add_argument("--patience", type=int, default=3) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--learning-rate", type=float, default=3e-4) parser.add_argument("--weight-decay", type=float, default=1e-3) parser.add_argument("--explanation-batch-size", type=int, default=32) args = parser.parse_args() run(args) if __name__ == "__main__": main()