建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,194 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user