File size: 25,153 Bytes
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
 
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
3c58630
 
 
 
 
 
 
 
 
 
 
eacee4b
 
 
3c58630
 
 
 
 
eacee4b
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eacee4b
3c58630
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
from __future__ import annotations

import argparse
import json
import os
import re
from pathlib import Path
from typing import Optional, Sequence

_MPLCONFIGDIR = Path(__file__).resolve().parents[2] / ".inference_work" / "matplotlib"
_MPLCONFIGDIR.mkdir(parents=True, exist_ok=True)
os.environ.setdefault("MPLCONFIGDIR", str(_MPLCONFIGDIR))

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.mplot3d.art3d import Poly3DCollection

try:
    from tqdm import tqdm
except Exception:
    def tqdm(iterable=None, *args, **kwargs):  # type: ignore[no-redef]
        return iterable if iterable is not None else ()

from official_demo_inference.paths import code_root as _code_root
from official_demo_inference.paths import default_demo_root as _default_demo_root
from official_demo_inference.paths import default_vertex_count_json as _default_vertex_count_json
from physformer.data.multiobj_utils_multiobj import load_mesh_vertex_counts, scene_info_from_metadata
from physformer.data.obj_io import load_obj_vertices_faces


EXCLUDED_DIR_NAMES = {"code", ".inference_work", "__pycache__"}
SPLIT_DIR_ALIASES = {
    "ood": "ood",
    "ood_examples": "ood",
    "indistribution": "indistribution",
}
COLORS = [
    (0.86, 0.24, 0.20, 1.0),
    (0.20, 0.64, 0.42, 1.0),
    (0.20, 0.44, 0.86, 1.0),
    (0.92, 0.67, 0.22, 1.0),
    (0.62, 0.32, 0.76, 1.0),
]
NAMED_COLORS = {
    "cow": (0.00, 0.62, 0.66, 1.0),
    "horse": (0.88, 0.30, 0.24, 1.0),
}
MESH_EDGE_COLOR = (0.05, 0.06, 0.07, 0.62)
LIGHT_DIRECTION = np.asarray([0.45, -0.65, 0.75], dtype=np.float32)
RIGID_RENDER_ALPHA = 0.96
ELASTIC_RENDER_ALPHA = 0.38


def _scalar_str(data: np.lib.npyio.NpzFile, key: str) -> Optional[str]:
    if key not in data.files:
        return None
    value = data[key]
    if isinstance(value, np.ndarray):
        if value.shape == ():
            return str(value.item())
        if value.size == 1:
            return str(value.reshape(-1)[0])
    return str(value)


def _sample_dir_from_vertices_path(vertices_path: Path) -> Path:
    if re.fullmatch(r"(?:gen|sample)_\d{2}", vertices_path.parent.name):
        return vertices_path.parent.parent
    return vertices_path.parent


def discover_vertices(root: Path, include: str, generations: set[int] | None) -> list[Path]:
    out: list[Path] = []
    for path in sorted(root.glob("**/vertices.npz")):
        try:
            rel = path.relative_to(root)
        except ValueError:
            continue
        parts = rel.parts
        if any(part in EXCLUDED_DIR_NAMES for part in parts):
            continue
        split = SPLIT_DIR_ALIASES.get(parts[0], "other") if parts else "other"
        if include != "all" and split != include:
            continue
        if generations is not None:
            parent_name = path.parent.name
            match = re.fullmatch(r"(?:gen|sample)_(\d{2})", parent_name)
            if match is None:
                if 0 not in generations:
                    continue
            else:
                gen_idx = int(match.group(1))
                if gen_idx not in generations:
                    continue
        out.append(path)
    return out


def _fixed_limits_from_metadata(meta: dict) -> tuple[tuple[float, float], tuple[float, float], tuple[float, float]]:
    bounds_min = np.asarray(meta.get("bounds_min", [-1.0, -1.0, -1.0]), dtype=np.float32)
    bounds_max = np.asarray(meta.get("bounds_max", [1.0, 1.0, 1.0]), dtype=np.float32)
    if bounds_min.shape != (3,) or bounds_max.shape != (3,):
        return ((-1.0, 1.0), (-1.0, 1.0), (-1.0, 1.0))
    return tuple((float(bounds_min[i]), float(bounds_max[i])) for i in range(3))  # type: ignore[return-value]


def _fixed_limits_from_cli(raw: str) -> Optional[tuple[tuple[float, float], tuple[float, float], tuple[float, float]]]:
    text = str(raw or "").strip()
    if not text:
        return None
    parts = [p for p in text.replace(";", ",").split(",") if p.strip()]
    if len(parts) != 6:
        raise ValueError(
            "--viz-fixed-limits must contain 6 comma-separated numbers: xmin,xmax,ymin,ymax,zmin,zmax; "
            f"got {raw!r}"
        )
    vals = [float(p.strip()) for p in parts]
    limits = ((vals[0], vals[1]), (vals[2], vals[3]), (vals[4], vals[5]))
    for lo, hi in limits:
        if not np.isfinite(lo) or not np.isfinite(hi) or hi <= lo:
            raise ValueError(f"Invalid --viz-fixed-limits range: {raw!r}")
    return limits


def _gt_frame_paths(sample_dir: Path) -> list[Path]:
    meshes_dir = sample_dir / "meshes"
    if not meshes_dir.is_dir():
        raise FileNotFoundError(f"Missing meshes directory: {meshes_dir}")
    out = sorted(meshes_dir.glob("combined_frame_*.obj"))
    if not out:
        raise FileNotFoundError(f"No combined_frame_*.obj files found under: {meshes_dir}")
    return out


def _load_gt_vertices(frame_paths: Sequence[Path], num_frames: int) -> np.ndarray:
    if len(frame_paths) < int(num_frames):
        raise ValueError(f"Need {num_frames} GT frames, found only {len(frame_paths)}")
    frames: list[np.ndarray] = []
    for path in frame_paths[: int(num_frames)]:
        verts, _ = load_obj_vertices_faces(str(path))
        frames.append(verts.astype(np.float32, copy=False))
    return np.stack(frames, axis=0).astype(np.float32, copy=False)


def _faces_by_object(first_obj_path: Path, vertex_slices: Sequence[tuple[int, int]], total_vertices: int) -> list[np.ndarray]:
    vertices, faces = load_obj_vertices_faces(str(first_obj_path))
    if int(vertices.shape[0]) != int(total_vertices):
        raise ValueError(
            f"Combined first-frame OBJ vertex count mismatch: obj has V={int(vertices.shape[0])} "
            f"but metadata sum is V={int(total_vertices)}. obj={first_obj_path}"
        )
    out: list[np.ndarray] = []
    for start, end in vertex_slices:
        start = int(start)
        end = int(end)
        in_range = (faces >= start) & (faces < end)
        keep = np.all(in_range, axis=1)
        obj_faces = faces[keep] - start
        if obj_faces.size == 0:
            raise ValueError(f"No faces found for object vertex slice ({start}, {end}) in {first_obj_path}")
        out.append(obj_faces.astype(np.int64))
    return out


def _shaded_facecolors(vertices: np.ndarray, faces: np.ndarray, base_color: tuple[float, float, float, float]) -> np.ndarray:
    tris = vertices[faces]
    normals = np.cross(tris[:, 1] - tris[:, 0], tris[:, 2] - tris[:, 0])
    normals /= np.maximum(np.linalg.norm(normals, axis=1, keepdims=True), 1e-8)
    light = LIGHT_DIRECTION / np.linalg.norm(LIGHT_DIRECTION)
    intensity = 0.42 + 0.58 * np.clip(normals @ light, 0.0, 1.0)
    base = np.asarray(base_color, dtype=np.float32)
    colors = np.empty((faces.shape[0], 4), dtype=np.float32)
    colors[:, :3] = np.clip(base[:3][None, :] * intensity[:, None] + 0.10 * (1.0 - intensity[:, None]), 0.0, 1.0)
    colors[:, 3] = base[3]
    return colors


def _color_for_object(index: int, object_name: str | None) -> tuple[float, float, float, float]:
    name = str(object_name or "").lower()
    for pattern, color in NAMED_COLORS.items():
        if pattern in name:
            return color
    return COLORS[int(index) % len(COLORS)]


def _with_alpha(color: tuple[float, float, float, float], alpha: float) -> tuple[float, float, float, float]:
    return (float(color[0]), float(color[1]), float(color[2]), float(alpha))


def _is_elastic_material_for_render(obj: object) -> bool:
    if not isinstance(obj, dict):
        return False
    material = obj.get("material")
    if isinstance(material, dict):
        kind = str(material.get("kind", "")).strip().lower()
        if kind in {"elastic", "soft"}:
            return True
        if kind in {"rigid", "hard"}:
            return False
        for key in ("effective_softness", "softness"):
            value = material.get(key)
            if isinstance(value, (int, float)):
                return float(value) >= 0.5
    for key in ("effective_softness", "softness"):
        value = obj.get(key)
        if isinstance(value, (int, float)):
            return float(value) >= 0.5
    return False


def _render_alphas_from_metadata(meta: dict, expected_count: int) -> list[float]:
    objects = meta.get("objects", [])
    if not isinstance(objects, list):
        objects = []
    out = [
        ELASTIC_RENDER_ALPHA if _is_elastic_material_for_render(obj) else RIGID_RENDER_ALPHA
        for obj in objects[: int(expected_count)]
    ]
    while len(out) < int(expected_count):
        out.append(RIGID_RENDER_ALPHA)
    return out


def _object_names_from_scene_json(sample_dir: Path, expected_count: int) -> list[str] | None:
    scene_path = sample_dir / "scene_multiobj.json"
    if not scene_path.is_file():
        return None
    try:
        with scene_path.open("r", encoding="utf-8") as handle:
            payload = json.load(handle)
    except Exception:
        return None
    objects = payload.get("objects") if isinstance(payload, dict) else None
    if not isinstance(objects, list) or len(objects) != int(expected_count):
        return None
    names: list[str] = []
    for obj in objects:
        if not isinstance(obj, dict):
            return None
        parts = [str(obj.get(key, "")) for key in ("name", "mesh_name", "mesh_path", "mesh_source", "mesh_used")]
        names.append(" ".join(part for part in parts if part.strip()))
    return names


def _render_multiobj_frame(
    vertices_by_obj: Sequence[np.ndarray],
    faces_by_obj: Sequence[np.ndarray],
    *,
    fixed_limits: tuple[tuple[float, float], tuple[float, float], tuple[float, float]],
    elev: float,
    azim: float,
    dpi: int,
    title: str = "",
    object_names: Sequence[str] | None = None,
    object_alphas: Sequence[float] | None = None,
) -> np.ndarray:
    fig = plt.figure(figsize=(5.4, 5.4), dpi=int(dpi), facecolor="#f7f8fb")
    ax = fig.add_subplot(1, 1, 1, projection="3d")
    ax.set_facecolor("#f7f8fb")

    for i, (vertices, faces) in enumerate(zip(vertices_by_obj, faces_by_obj)):
        vertices = np.asarray(vertices, dtype=np.float32)
        faces = np.asarray(faces, dtype=np.int64)
        if vertices.size == 0 or faces.size == 0:
            continue
        object_name = object_names[i] if object_names is not None and i < len(object_names) else None
        alpha = object_alphas[i] if object_alphas is not None and i < len(object_alphas) else RIGID_RENDER_ALPHA
        color = _with_alpha(_color_for_object(i, object_name), alpha)
        facecolors = _shaded_facecolors(vertices, faces, color)
        poly = Poly3DCollection(
            vertices[faces],
            facecolors=facecolors,
            edgecolors=MESH_EDGE_COLOR,
            linewidths=0.28,
            alpha=alpha,
            antialiased=True,
        )
        ax.add_collection3d(poly)

    (x_min, x_max), (y_min, y_max), (z_min, z_max) = fixed_limits
    ax.set_xlim(x_min, x_max)
    ax.set_ylim(y_min, y_max)
    ax.set_zlim(z_min, z_max)
    ax.set_box_aspect([1, 1, 1])
    ax.view_init(elev=float(elev), azim=float(azim))
    try:
        ax.set_proj_type("persp", focal_length=0.85)
    except TypeError:
        ax.set_proj_type("persp")
    ax.set_xlabel("")
    ax.set_ylabel("")
    ax.set_zlabel("")
    ax.tick_params(axis="both", which="major", labelsize=7, colors="#667085", pad=1)
    ax.grid(True, linestyle="-", linewidth=0.45, alpha=0.22)
    for axis in (ax.xaxis, ax.yaxis, ax.zaxis):
        axis.pane.set_facecolor((0.95, 0.96, 0.98, 0.72))
        axis.pane.set_edgecolor((0.78, 0.81, 0.86, 0.45))
    if title:
        ax.set_title(str(title), fontsize=13, fontweight="bold", color="#111827", pad=10)

    fig.subplots_adjust(left=0.02, right=0.98, bottom=0.10, top=0.93 if title else 0.99)
    fig.canvas.draw()
    width, height = fig.canvas.get_width_height()
    image = np.frombuffer(fig.canvas.buffer_rgba(), dtype=np.uint8).reshape(height, width, 4)[:, :, :3]
    plt.close(fig)
    return image


def _compose_side_by_side(
    left: np.ndarray,
    right: np.ndarray,
    *,
    sample_title: str,
    subset_label: str,
    dpi: int,
) -> np.ndarray:
    fig, axes = plt.subplots(1, 2, figsize=(10, 5), dpi=int(dpi))
    title = str(sample_title).strip()
    subset = str(subset_label).strip()
    if subset:
        title = f"{subset} | {title}" if title else subset
    if title:
        fig.suptitle(title, fontsize=13, fontweight="bold")
    for ax, image, panel_title in zip(axes, [left, right], ["Ground Truth", "Inference"]):
        ax.imshow(image)
        ax.set_title(panel_title, fontsize=11)
        ax.axis("off")
    fig.tight_layout(rect=[0.0, 0.0, 1.0, 0.95] if title else None)
    fig.canvas.draw()
    width, height = fig.canvas.get_width_height()
    out = np.frombuffer(fig.canvas.buffer_rgba(), dtype=np.uint8).reshape(height, width, 4)[:, :, :3]
    plt.close(fig)
    return out


def _save_animation(frames: Sequence[np.ndarray], *, out_gif: Optional[Path], out_mp4: Optional[Path], fps: int) -> None:
    if out_gif is None and out_mp4 is None:
        return
    try:
        import imageio.v2 as imageio  # type: ignore
    except Exception as exc:
        raise RuntimeError("Saving GIF/MP4 requires imageio. Install with: pip install imageio imageio-ffmpeg") from exc

    if out_gif is not None:
        imageio.mimsave(str(out_gif), list(frames), duration=1.0 / max(1, int(fps)), loop=0)
    if out_mp4 is not None:
        try:
            with imageio.get_writer(str(out_mp4), fps=max(1, int(fps)), codec="libx264", quality=8) as writer:
                for frame in frames:
                    writer.append_data(frame)
        except Exception as exc:
            raise RuntimeError("MP4 saving failed. You likely need ffmpeg support: pip install imageio-ffmpeg") from exc


def _paths_from_npz(vertices_path: Path, data: np.lib.npyio.NpzFile) -> tuple[Path, Path, Path]:
    sample_dir = _sample_dir_from_vertices_path(vertices_path)
    cond_sample_dir = Path(_scalar_str(data, "cond_sample_dir") or sample_dir)
    metadata_path = Path(_scalar_str(data, "cond_metadata_path") or (cond_sample_dir / "metadata.json"))

    if not cond_sample_dir.is_dir():
        cond_sample_dir = sample_dir
    if not metadata_path.is_file():
        metadata_path = cond_sample_dir / "metadata.json"
    if not metadata_path.is_file():
        metadata_path = sample_dir / "metadata.json"
    if not metadata_path.is_file():
        raise FileNotFoundError(f"Could not find metadata for {vertices_path}")

    first_obj_path = cond_sample_dir / "meshes" / "combined_frame_000.obj"
    if not first_obj_path.is_file():
        first_obj_path = sample_dir / "meshes" / "combined_frame_000.obj"
    if not first_obj_path.is_file():
        raise FileNotFoundError(f"Could not find first-frame combined OBJ for {vertices_path}")
    return cond_sample_dir, metadata_path, first_obj_path


def render_vertices_file(vertices_path: Path, args: argparse.Namespace, vertex_counts: dict[str, int]) -> bool:
    gen_dir = vertices_path.parent
    pred_gif = gen_dir / "inference.gif" if args.save_gif else None
    pred_mp4 = gen_dir / "inference.mp4" if args.save_mp4 else None
    gt_gif = gen_dir / "GT.gif" if args.save_gt_gif else None
    gt_mp4 = gen_dir / "GT.mp4" if args.save_gt_mp4 else None
    compare_base = Path(str(args.compare_out_name)).stem
    compare_gif = gen_dir / f"{compare_base}.gif" if args.save_compare_gif else None
    compare_mp4 = gen_dir / f"{compare_base}.mp4" if args.save_compare_mp4 else None
    requested = [p for p in [pred_gif, pred_mp4, gt_gif, gt_mp4, compare_gif, compare_mp4] if p is not None]
    if requested and all(p.is_file() for p in requested) and not bool(args.overwrite):
        print(f"[SKIP] renders exist: {gen_dir}", flush=True)
        return False

    with np.load(vertices_path, allow_pickle=False) as data:
        if "vertices" not in data.files:
            raise KeyError(f"{vertices_path} does not contain a 'vertices' array")
        vertices = np.asarray(data["vertices"], dtype=np.float32)
        cond_sample_dir, metadata_path, first_obj_path = _paths_from_npz(vertices_path, data)

    if vertices.ndim != 3 or vertices.shape[-1] != 3:
        raise ValueError(f"Expected vertices shape (F,V,3), got {tuple(vertices.shape)} in {vertices_path}")

    with metadata_path.open("r", encoding="utf-8") as handle:
        meta = json.load(handle)
    scene = scene_info_from_metadata(str(metadata_path), vertex_counts=vertex_counts, max_num_objects=int(args.max_num_objects))
    if int(vertices.shape[1]) < int(scene.total_vertices):
        raise ValueError(
            f"{vertices_path} has V={int(vertices.shape[1])}, but metadata expects V={int(scene.total_vertices)}"
        )

    pred_vertices = vertices[:, : int(scene.total_vertices), :].astype(np.float32, copy=False)
    faces_by_obj = _faces_by_object(first_obj_path, scene.vertex_slices, int(scene.total_vertices))
    sample_dir = _sample_dir_from_vertices_path(vertices_path)
    object_names = _object_names_from_scene_json(sample_dir, len(scene.vertex_slices))
    if object_names is None:
        object_names = [
            f"{str(name)} {str(path)}"
            for name, path in zip(scene.mesh_names, scene.mesh_paths)
        ]
    object_alphas = _render_alphas_from_metadata(meta, len(scene.vertex_slices))
    fixed_limits = _fixed_limits_from_cli(str(args.viz_fixed_limits)) or _fixed_limits_from_metadata(meta)

    gt_vertices = None
    if args.save_gt_gif or args.save_gt_mp4 or args.save_compare_gif or args.save_compare_mp4:
        gt_vertices = _load_gt_vertices(_gt_frame_paths(cond_sample_dir), int(pred_vertices.shape[0]))
        if int(gt_vertices.shape[1]) != int(scene.total_vertices):
            raise ValueError(
                f"GT vertices have V={int(gt_vertices.shape[1])}, but metadata expects V={int(scene.total_vertices)}. "
                f"sample_dir={cond_sample_dir}"
            )

    pred_frames: list[np.ndarray] = []
    gt_frames: list[np.ndarray] = []
    compare_frames: list[np.ndarray] = []
    frame_range = range(int(pred_vertices.shape[0]))
    if not bool(args.verbose):
        frame_range = tqdm(frame_range, desc=f"render[{vertices_path.parent.name}]", unit="frame", leave=False)

    sample_title = str(sample_dir.relative_to(args.demo_root))
    for frame_idx in frame_range:
        pred_by_obj = [pred_vertices[frame_idx, start:end, :] for start, end in scene.vertex_slices]
        pred_img = _render_multiobj_frame(
            pred_by_obj,
            faces_by_obj,
            fixed_limits=fixed_limits,
            elev=float(args.viz_elev),
            azim=float(args.viz_azim),
            dpi=int(args.compare_render_dpi if (args.save_compare_gif or args.save_compare_mp4) else args.render_dpi),
            title="inference" if (args.save_gif or args.save_mp4) else "",
            object_names=object_names,
            object_alphas=object_alphas,
        )
        if args.save_gif or args.save_mp4:
            pred_frames.append(pred_img)
        if args.save_gt_gif or args.save_gt_mp4 or args.save_compare_gif or args.save_compare_mp4:
            assert gt_vertices is not None
            gt_by_obj = [gt_vertices[frame_idx, start:end, :] for start, end in scene.vertex_slices]
            gt_img = _render_multiobj_frame(
                gt_by_obj,
                faces_by_obj,
                fixed_limits=fixed_limits,
                elev=float(args.viz_elev),
                azim=float(args.viz_azim),
                dpi=int(args.compare_render_dpi),
                title="GT" if (args.save_gt_gif or args.save_gt_mp4) else "",
                object_names=object_names,
                object_alphas=object_alphas,
            )
            if args.save_gt_gif or args.save_gt_mp4:
                gt_frames.append(gt_img)
        if args.save_compare_gif or args.save_compare_mp4:
            compare_frames.append(
                _compose_side_by_side(
                    gt_img,
                    pred_img,
                    sample_title=sample_title,
                    subset_label=str(args.compare_subset_label),
                    dpi=int(args.compare_compose_dpi),
                )
            )

    if args.save_gif or args.save_mp4:
        _save_animation(
            pred_frames,
            out_gif=pred_gif if (pred_gif is not None and (args.overwrite or not pred_gif.exists())) else None,
            out_mp4=pred_mp4 if (pred_mp4 is not None and (args.overwrite or not pred_mp4.exists())) else None,
            fps=int(args.fps),
        )
    if args.save_gt_gif or args.save_gt_mp4:
        _save_animation(
            gt_frames,
            out_gif=gt_gif if (gt_gif is not None and (args.overwrite or not gt_gif.exists())) else None,
            out_mp4=gt_mp4 if (gt_mp4 is not None and (args.overwrite or not gt_mp4.exists())) else None,
            fps=int(args.fps),
        )
    if args.save_compare_gif or args.save_compare_mp4:
        _save_animation(
            compare_frames,
            out_gif=compare_gif if (compare_gif is not None and (args.overwrite or not compare_gif.exists())) else None,
            out_mp4=compare_mp4 if (compare_mp4 is not None and (args.overwrite or not compare_mp4.exists())) else None,
            fps=int(args.fps),
        )
    print(f"[RENDERED] {vertices_path}", flush=True)
    return True


def build_argparser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        description="Render existing official-demo vertices.npz files without running inference.",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )
    parser.add_argument("--demo-root", type=Path, default=_default_demo_root())
    parser.add_argument("--mesh-vertex-count-json", type=Path, default=_default_vertex_count_json())
    parser.add_argument("--include", choices=["all", "ood", "indistribution"], default="all")
    parser.add_argument("--generation", type=int, action="append", default=None, help="Only render a rollout sample index, e.g. 0 for sample_00. Can be repeated.")
    parser.add_argument("--max-samples", type=int, default=0, help="Limit the number of vertices.npz files rendered; 0 means no limit.")
    parser.add_argument("--max-num-objects", type=int, default=10)
    parser.add_argument("--save-gif", action="store_true")
    parser.add_argument("--save-mp4", action="store_true")
    parser.add_argument("--save-gt-gif", action="store_true")
    parser.add_argument("--save-gt-mp4", action="store_true")
    parser.add_argument("--save-compare-gif", action="store_true")
    parser.add_argument("--save-compare-mp4", action="store_true")
    parser.add_argument("--compare-out-name", default="traj_compare_gt_vs_infer")
    parser.add_argument("--compare-subset-label", default="")
    parser.add_argument("--fps", type=int, default=12)
    parser.add_argument("--render-dpi", type=int, default=160)
    parser.add_argument("--compare-render-dpi", type=int, default=160)
    parser.add_argument("--compare-compose-dpi", type=int, default=160)
    parser.add_argument("--viz-elev", type=float, default=30.0)
    parser.add_argument("--viz-azim", type=float, default=-45.0)
    parser.add_argument("--viz-fixed-limits", default="")
    parser.add_argument("--overwrite", action=argparse.BooleanOptionalAction, default=False)
    parser.add_argument("--verbose", action="store_true")
    return parser


def main(argv: Sequence[str] | None = None) -> int:
    parser = build_argparser()
    args = parser.parse_args(argv)
    args.demo_root = args.demo_root.expanduser().resolve()
    args.mesh_vertex_count_json = args.mesh_vertex_count_json.expanduser().resolve()

    if not args.demo_root.is_dir():
        parser.error(f"--demo-root is not a directory: {args.demo_root}")
    if not args.mesh_vertex_count_json.is_file():
        parser.error(f"--mesh-vertex-count-json is not a file: {args.mesh_vertex_count_json}")
    if not (args.save_gif or args.save_mp4 or args.save_gt_gif or args.save_gt_mp4 or args.save_compare_gif or args.save_compare_mp4):
        parser.error("Choose at least one render output: --save-mp4, --save-gif, --save-gt-mp4, --save-gt-gif, --save-compare-mp4, or --save-compare-gif")

    generations = set(args.generation) if args.generation is not None else None
    vertex_paths = discover_vertices(args.demo_root, args.include, generations)
    if args.max_samples and int(args.max_samples) > 0:
        vertex_paths = vertex_paths[: int(args.max_samples)]
    if not vertex_paths:
        parser.error(f"No vertices.npz files found under {args.demo_root}")

    vertex_counts = load_mesh_vertex_counts(str(args.mesh_vertex_count_json))

    rendered = 0
    iterator = vertex_paths if args.verbose else tqdm(vertex_paths, desc="outputs", unit="npz")
    for vertices_path in iterator:
        rendered += int(render_vertices_file(vertices_path, args, vertex_counts))
    print(f"[done] rendered {rendered} of {len(vertex_paths)} vertices.npz files", flush=True)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())