File size: 5,626 Bytes
eafbe80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
#!/usr/bin/env python3
from __future__ import annotations

import argparse
import csv
import json
import os
from pathlib import Path


IMG_EXTS = {".png", ".jpg", ".jpeg", ".webp"}


def _step_id(path: Path) -> int:
    stem = path.stem
    if stem.startswith("Step-"):
        try:
            return int(stem.split("-", 1)[1])
        except Exception:
            return -1
    return -1


def find_ckpts(outputs_root: Path, include_dyn: bool, ckpt_policy: str = "epoch0") -> list[Path]:
    """Return checkpoints under outputs_root according to ckpt_policy.

    Evaluation candidacy is intentionally file-based: if a run directory has
    a matching checkpoint, it is considered finished enough to evaluate. The
    include_dyn flag is kept for backward compatibility and does not filter
    Echo-Memory outputs by default.
    """
    policy = (ckpt_policy or "epoch0").strip().lower()
    if policy == "epoch0":
        return sorted(outputs_root.rglob("epoch-0.safetensors"))
    if policy == "all_steps":
        return sorted(outputs_root.rglob("Step-*.safetensors"), key=lambda p: (str(p.parent), _step_id(p)))
    if policy == "latest":
        by_dir: dict[Path, Path] = {}
        for p in sorted(outputs_root.rglob("*.safetensors")):
            if p.name == "epoch-0.safetensors":
                by_dir.setdefault(p.parent, p)
            elif p.name.startswith("Step-"):
                prev = by_dir.get(p.parent)
                if prev is None or (prev.name == "epoch-0.safetensors") or _step_id(p) > _step_id(prev):
                    by_dir[p.parent] = p
        return sorted(by_dir.values())
    raise ValueError(f"unknown ckpt_policy={ckpt_policy}; expected epoch0/latest/all_steps")


def read_train_samples(metadata: Path, dataset_base: Path, limit: int) -> list[dict]:
    rows = []
    with metadata.open("r", encoding="utf-8", newline="") as f:
        for row in csv.DictReader(f):
            video_name = (row.get("video_name") or "").strip()
            if not video_name:
                continue
            try:
                start_frame = int(row.get("start_frame", 0) or 0)
            except (TypeError, ValueError):
                start_frame = 0
            try:
                end_frame = int(row.get("end_frame", start_frame + 80) or (start_frame + 80))
            except (TypeError, ValueError):
                end_frame = start_frame + 80
            frame_path = dataset_base / "frames" / video_name / f"{start_frame:04d}.png"
            if not frame_path.is_file():
                continue
            rows.append(
                {
                    "sample_id": f"train_{len(rows):05d}_{video_name.replace('/', '__')}_{start_frame:06d}",
                    "domain": "train",
                    "video_name": video_name,
                    "start_frame": start_frame,
                    "end_frame": end_frame,
                    "prompt": row.get("prompt") or "A scene.",
                    "first_frame_image": str(frame_path),
                    "dataset_base": str(dataset_base),
                }
            )
            if limit and len(rows) >= limit:
                break
    return rows


def read_ood_samples(ood_dir: Path, prompt: str, limit: int) -> list[dict]:
    imgs = [p for p in sorted(ood_dir.iterdir()) if p.suffix.lower() in IMG_EXTS]
    if limit:
        imgs = imgs[:limit]
    return [
        {
            "sample_id": f"ood_{i:05d}_{p.stem}",
            "domain": "ood",
            "video_name": None,
            "start_frame": 0,
            "end_frame": 80,
            "prompt": prompt,
            "first_frame_image": str(p),
            "dataset_base": None,
        }
        for i, p in enumerate(imgs)
    ]


def main() -> None:
    repo_root = Path(__file__).resolve().parents[3]
    default_dataset = repo_root / "data" / "Context-as-Memory-Dataset"
    ap = argparse.ArgumentParser()
    ap.add_argument("--outputs-root", default=str(repo_root / "outputs"))
    ap.add_argument("--dataset-base", default=str(default_dataset))
    ap.add_argument("--metadata", default=str(default_dataset / "metadata_full.csv"))
    ap.add_argument("--ood-dir", default=str(repo_root / "assets" / "opendomain_revisit"))
    ap.add_argument("--out", required=True)
    ap.add_argument("--train-limit", type=int, default=8)
    ap.add_argument("--ood-limit", type=int, default=8)
    ap.add_argument("--ood-prompt", default="A toy bear in the same static scene. Preserve the bear appearance and the scene layout after camera revisit.")
    ap.add_argument("--include-dynmembench", action="store_true")
    ap.add_argument("--ckpt-policy", default="epoch0", choices=["epoch0", "latest", "all_steps"])
    args = ap.parse_args()

    ckpts = find_ckpts(Path(args.outputs_root), args.include_dynmembench, args.ckpt_policy)
    train_samples = read_train_samples(Path(args.metadata), Path(args.dataset_base), args.train_limit)
    ood_samples = read_ood_samples(Path(args.ood_dir), args.ood_prompt, args.ood_limit)
    payload = {
        "ckpts": [{"ckpt": str(p), "run_id": p.parent.name} for p in ckpts],
        "samples": train_samples + ood_samples,
        "dataset_base": args.dataset_base,
        "metadata": args.metadata,
        "ood_dir": args.ood_dir,
        "ckpt_policy": args.ckpt_policy,
    }
    out = Path(args.out)
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
    print(f"[prepare_eval_manifest] ckpts={len(ckpts)} train={len(train_samples)} ood={len(ood_samples)} -> {out}")


if __name__ == "__main__":
    main()