Add Q3 MoFE router visualizations and explanations
This commit is contained in:
@@ -1,569 +1,6 @@
|
||||
"""Train Q3 on the official training split and explain Attachment 4 cases."""
|
||||
from __future__ import annotations
|
||||
"""Compatibility entry point for the first Q3 explanation experiment."""
|
||||
|
||||
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)
|
||||
from .run_experiments import main
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user