Files
modeling_zhaocui/final/q3/train_interpretable.py
T

571 lines
29 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()