195 lines
8.2 KiB
Python
195 lines
8.2 KiB
Python
"""Summarize the saved D0-D3 alignment-diagnostic metrics without retraining."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import csv
|
||
from collections import defaultdict
|
||
from pathlib import Path
|
||
import shutil
|
||
from statistics import mean
|
||
from typing import Any
|
||
|
||
|
||
METRICS = (
|
||
"mvr",
|
||
"normalized_entropy",
|
||
"c_row",
|
||
"trajectory_span",
|
||
"mean_absolute_time_center_error",
|
||
"gaussian_target_kl",
|
||
)
|
||
|
||
|
||
def merge_trial_metrics(run_dir: Path, destination: Path) -> None:
|
||
"""Rebuild the combined file with a variant label for PE controls."""
|
||
rows: list[dict[str, str]] = []
|
||
fields: list[str] | None = None
|
||
for experiment in ("D0", "D1", "D2", "D3"):
|
||
for path in sorted((run_dir / experiment).glob("*_metrics.csv")):
|
||
variant = path.name.removesuffix("_metrics.csv")
|
||
with path.open(newline="", encoding="utf-8-sig") as handle:
|
||
reader = csv.DictReader(handle)
|
||
if fields is None:
|
||
fields = ["experiment", "method", "variant"] + [
|
||
name for name in (reader.fieldnames or [])
|
||
if name not in {"experiment", "method", "variant"}
|
||
]
|
||
for row in reader:
|
||
row["variant"] = variant
|
||
rows.append(row)
|
||
if not fields:
|
||
raise FileNotFoundError(f"no per-trial metric CSV files found under {run_dir}")
|
||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||
with destination.open("w", newline="", encoding="utf-8-sig") as handle:
|
||
writer = csv.DictWriter(handle, fieldnames=fields)
|
||
writer.writeheader()
|
||
writer.writerows(rows)
|
||
|
||
|
||
def summarize(source: Path, destination: Path) -> list[dict[str, Any]]:
|
||
groups: dict[tuple[str, str, str], list[dict[str, str]]] = defaultdict(list)
|
||
with source.open(newline="", encoding="utf-8-sig") as handle:
|
||
for row in csv.DictReader(handle):
|
||
if row["experiment"] in {"D2", "D3"}:
|
||
groups[(row["experiment"], row["method"], row["modality"])].append(row)
|
||
|
||
summaries: list[dict[str, Any]] = []
|
||
for (experiment, method, modality), rows in sorted(groups.items()):
|
||
summary: dict[str, Any] = {
|
||
"experiment": experiment,
|
||
"method": method,
|
||
"modality": modality,
|
||
"sample_count": len(rows),
|
||
}
|
||
for metric in METRICS:
|
||
values = [float(row[metric]) for row in rows if row.get(metric, "") != ""]
|
||
summary[f"mean_{metric}"] = mean(values) if values else ""
|
||
summary[f"n_{metric}"] = len(values)
|
||
summaries.append(summary)
|
||
|
||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||
with destination.open("w", newline="", encoding="utf-8-sig") as handle:
|
||
writer = csv.DictWriter(handle, fieldnames=list(summaries[0]))
|
||
writer.writeheader()
|
||
writer.writerows(summaries)
|
||
return summaries
|
||
|
||
|
||
def build_report_bundle(run_dir: Path) -> Path:
|
||
"""Copy compact diagnostic evidence, leaving large checkpoints local."""
|
||
bundle = run_dir / "report_bundle"
|
||
bundle.mkdir(parents=True, exist_ok=True)
|
||
for name in (
|
||
"debug_summary.csv",
|
||
"metric_summary.csv",
|
||
"per_sample_metrics.csv",
|
||
"run_manifest.json",
|
||
"experiment.log",
|
||
):
|
||
source = run_dir / name
|
||
if source.exists():
|
||
shutil.copy2(source, bundle / name)
|
||
for experiment in ("D0", "D1", "D2", "D3"):
|
||
source_dir = run_dir / experiment
|
||
target_dir = bundle / experiment
|
||
target_dir.mkdir(exist_ok=True)
|
||
for pattern in ("*_history.csv", "*_metrics.csv", "*_heatmap.png", "*_trajectory.png"):
|
||
for source in source_dir.glob(pattern):
|
||
shutil.copy2(source, target_dir / source.name)
|
||
synthetic_source = run_dir / "synthetic"
|
||
synthetic_target = bundle / "synthetic"
|
||
synthetic_target.mkdir(exist_ok=True)
|
||
for name in ("metrics.csv", "training_history.csv", "run_manifest.json"):
|
||
source = synthetic_source / name
|
||
if source.exists():
|
||
shutil.copy2(source, synthetic_target / name)
|
||
for source in synthetic_source.rglob("*.png"):
|
||
target = synthetic_target / source.relative_to(synthetic_source)
|
||
target.parent.mkdir(parents=True, exist_ok=True)
|
||
shutil.copy2(source, target)
|
||
source_time_dir = run_dir / "source_time"
|
||
source_time_target = bundle / "source_time"
|
||
source_time_target.mkdir(exist_ok=True)
|
||
for name in ("summary.csv", "per_sample_metrics.csv", "run_manifest.json"):
|
||
source = source_time_dir / name
|
||
if source.exists():
|
||
shutil.copy2(source, source_time_target / name)
|
||
for variant_dir in source_time_dir.glob("M*"):
|
||
target_dir = source_time_target / variant_dir.name
|
||
target_dir.mkdir(exist_ok=True)
|
||
for pattern in ("history.csv", "metrics.csv", "*_heatmap.png", "*_trajectory.png"):
|
||
for source in variant_dir.glob(pattern):
|
||
shutil.copy2(source, target_dir / source.name)
|
||
heldout_dir = run_dir / "heldout"
|
||
heldout_target = bundle / "heldout"
|
||
heldout_target.mkdir(exist_ok=True)
|
||
for name in (
|
||
"heldout_summary.csv",
|
||
"per_sample_metrics.csv",
|
||
"training_summary.csv",
|
||
"training_history.csv",
|
||
"run_manifest.json",
|
||
):
|
||
source = heldout_dir / name
|
||
if source.exists():
|
||
shutil.copy2(source, heldout_target / name)
|
||
for variant_dir in heldout_dir.glob("M*"):
|
||
target_dir = heldout_target / variant_dir.name
|
||
target_dir.mkdir(exist_ok=True)
|
||
for pattern in ("history.csv", "heldout_metrics.csv", "*_heatmap.png", "*_trajectory.png"):
|
||
for source in variant_dir.glob(pattern):
|
||
shutil.copy2(source, target_dir / source.name)
|
||
readme = """# Q1 M3/M4 可学习性诊断结果包
|
||
|
||
本包包含 D0–D3 诊断的汇总表、运行清单、日志、各试验的损失历史、逐样本指标和代表性图像。PyTorch 检查点与逐样本 `.npz` 注意力矩阵留在上级 `outputs/alignment_debug/`,因此结果包较轻。
|
||
|
||
- D0/D1:单样本过拟合诊断。
|
||
- D1-S:输入仅为合成时间坐标;M3/M4 都能把 Gaussian 目标拟合至约 1e-4 KL。
|
||
- D2/D3:100 条样本上的同集训练/评价诊断,单个随机种子,不用于声称泛化。
|
||
- D4:真实单样本加入显式时间 key/query 特征后,Audio 注意力形成局部时间带。
|
||
- D5:按 `video_id` 留出 20 条样本;M4 的 Audio/Vision 时间带能迁移,M3 有改善但仍未充分贴近目标。
|
||
- Gaussian 时间目标是根据源时间戳生成的弱先验,不是人工对齐真值。
|
||
- M3 的 Gaussian KL 取 Audio/Vision 平均,M4 取 Text/Audio/Vision 平均;KL 数值仅用于各自的优化诊断,不可作为方法排名。
|
||
- D1 真实特征上 Text/Vision 能拟合,Audio 仍失败;合成对照成功,提示真实源特征的显式时间身份值得优先验证。
|
||
- 全数据 Q/K 梯度仍非零,但学习到的注意力没有稳定形成局部时间带;加入重构和对比目标后 Audio 塌缩更明显。
|
||
- 环境:Fedora WSL,`uv`,NVIDIA GeForce RTX 5070 Ti,Python 3.14.7,PyTorch 2.14.0+cu130。
|
||
|
||
详细解释见项目根目录的 `RESULTS.md`。
|
||
"""
|
||
(bundle / "README.md").write_text(readme, encoding="utf-8")
|
||
return bundle
|
||
|
||
|
||
def main() -> None:
|
||
project = Path(__file__).resolve().parents[1]
|
||
parser = argparse.ArgumentParser(description=__doc__)
|
||
parser.add_argument(
|
||
"--input",
|
||
type=Path,
|
||
default=project / "outputs/alignment_debug/per_sample_metrics.csv",
|
||
)
|
||
parser.add_argument(
|
||
"--output",
|
||
type=Path,
|
||
default=project / "outputs/alignment_debug/metric_summary.csv",
|
||
)
|
||
args = parser.parse_args()
|
||
merge_trial_metrics(args.input.parent, args.input)
|
||
for row in summarize(args.input, args.output):
|
||
values = " ".join(
|
||
f"{name}={row[f'mean_{name}']:.4f}"
|
||
for name in METRICS
|
||
if row[f"mean_{name}"] != ""
|
||
)
|
||
print(
|
||
f"{row['experiment']} {row['method']} {row['modality']} "
|
||
f"n={row['sample_count']} {values}"
|
||
)
|
||
print(f"Wrote {args.output}")
|
||
print(f"Built report bundle at {build_report_bundle(args.input.parent)}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|