fishxinyu's picture
download
raw
21.6 kB
"""Evaluate depth consistency of generated videos against condition depth videos.
Pipeline:
1. Load condition depth videos from <dataset-dir>/depth/<id>.mp4 (per-video
normalized depth stored as grayscale, decoded here to the 0-255 scale)
2. Extract depth from generated videos using Video-Depth-Anything, per-video
normalized to the same 0-255 scale
3. Optionally load a generated-person exclusion mask
4. Align generated depth to condition depth via least-squares scale+shift in
disparity space using only unmasked background pixels
5. Compute depth metrics (see eval_metrics below) over those same pixels, in
0-255 grayscale units - relative disparity, not metric depth
Usage:
python eval_depth/eval.py \
--generated-dir /home/xinyuy/dataset_processing/flexcombine-bench/results/vacebench/ltx-2/depth \
--dataset-dir /home/xinyuy/flexvideo/VACE/benchmarks/VACE-Benchmark/depth \
--depth-repo-path /home/xinyuy/dataset_processing/flexcombine-bench/tools/Video-Depth-Anything \
--output-dir /home/xinyuy/dataset_processing/flexcombine-bench/results/vacebench/ltx-2/depth/control
"""
import argparse
import csv
import json
import math
import sys
from fractions import Fraction
from pathlib import Path
import av
import numpy as np
import torch
import torch.nn.functional as F
from tqdm import tqdm
import metric
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from derange import deranged_pairing # noqa: E402
device ="cuda" if torch.cuda.is_available() else "cpu"
eval_metrics = [
"mae",
"rmse_linear",
]
# ---------------------------------------------------------------------------
# Video I/O
# ---------------------------------------------------------------------------
def _read_video(video_path: str | Path) -> tuple[torch.Tensor, float]:
"""Read a video file into a [T, C, H, W] float tensor in [0, 1] range."""
with av.open(str(video_path)) as container:
stream = container.streams.video[0]
fps = float(stream.average_rate or stream.base_rate or 24)
frames = [f.to_ndarray(format="rgb24") for f in container.decode(video=0)]
frames_np = np.stack(frames, axis=0) # [T, H, W, C] uint8
video = torch.from_numpy(frames_np).float().div(255.0)
return video.permute(0, 3, 1, 2), fps # [T, C, H, W]
def _resize_and_center_crop(video: torch.Tensor, height: int, width: int) -> torch.Tensor:
"""Resize-to-fill then center-crop — mirrors ic_lora.py::resize_and_center_crop."""
_, _, src_h, src_w = video.shape
scale = max(height / src_h, width / src_w)
new_h = math.ceil(src_h * scale)
new_w = math.ceil(src_w * scale)
video = F.interpolate(video, size=(new_h, new_w), mode="bilinear", align_corners=False)
crop_top = (new_h - height) // 2
crop_left = (new_w - width) // 2
return video[:, :, crop_top: crop_top + height, crop_left: crop_left + width]
# ---------------------------------------------------------------------------
# Depth model
# ---------------------------------------------------------------------------
_DEPTH_MODEL_CONFIGS = {
"vits": {"encoder": "vits", "features": 64, "out_channels": [48, 96, 192, 384]},
"vitb": {"encoder": "vitb", "features": 128, "out_channels": [96, 192, 384, 768]},
"vitl": {"encoder": "vitl", "features": 256, "out_channels": [256, 512, 1024, 1024]},
}
def _load_depth_model(encoder: str, device: str, repo_path: Path):
"""Load VideoDepthAnything from a local repo checkout."""
repo_root = str(repo_path.resolve())
if repo_root not in sys.path:
sys.path.insert(0, repo_root)
try:
from video_depth_anything.video_depth import VideoDepthAnything
except ImportError as exc:
raise ImportError(
f"Cannot import video_depth_anything from {repo_root}.\n"
"Make sure --depth-repo-path points to the root of a cloned "
"Video-Depth-Anything repo with its dependencies installed."
) from exc
ckpt = repo_path.resolve() / "checkpoints" / f"video_depth_anything_{encoder}.pth"
if not ckpt.exists():
raise FileNotFoundError(
f"Checkpoint not found: {ckpt}\n"
"Download from https://huggingface.co/depth-anything/Video-Depth-Anything and "
f"place it at {ckpt}"
)
from video_depth_anything.video_depth import VideoDepthAnything # noqa: F811
model = VideoDepthAnything(**_DEPTH_MODEL_CONFIGS[encoder])
model.load_state_dict(torch.load(ckpt, map_location="cpu"), strict=True)
return model.to(device).eval()
def _extract_depth(
video: torch.Tensor,
fps: float,
model,
device: str,
input_size: int = 518,
) -> np.ndarray:
"""Run Video-Depth-Anything and normalize per-video to [0, 1], then rescale
to the 0-255 grayscale range the metrics are computed in.
Mirrors compute_reference.py::compute_depth_reference so the output is
comparable to stored condition depth videos.
Returns:
depths: [T, H, W] float numpy array in [0, 255]
"""
frames_np = (video.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8)
with torch.inference_mode():
depths, _ = model.infer_video_depth(frames_np, fps, input_size=input_size, device=device)
d_min, d_max = float(depths.min()), float(depths.max())
if d_max > d_min:
depths = (depths - d_min) / (d_max - d_min)
else:
depths = np.zeros_like(depths)
return (depths * 255.0).astype(np.float32) # [T, H, W] in [0, 255]
def _load_condition_depth(
cond_path: str | Path,
target_height: int,
target_width: int,
) -> tuple[np.ndarray, float]:
"""Load a stored condition depth video, cropped to the generated video's resolution.
Applies the same resize-to-fill + center-crop that ic_lora.py uses so we compare
the same spatial region the model conditioned on.
Returns:
depths: [T, H, W] float array in [0, 255]
fps: video frame rate
"""
video, fps = _read_video(cond_path) # [T, C, H, W] in [0, 1]
video = _resize_and_center_crop(video, target_height, target_width)
return (video[:, 0, :, :].numpy() * 255.0).astype(np.float32), fps # [T, H, W] in [0, 255]
# ---------------------------------------------------------------------------
# Alignment and metrics (adapted from eval_depthcrafter)
# ---------------------------------------------------------------------------
def _align_depth(
gen_depth: np.ndarray,
cond_depth: np.ndarray,
valid_mask: np.ndarray | None = None,
) -> np.ndarray:
"""Least-squares scale+shift alignment of gen_depth to cond_depth.
Both inputs are [T, H, W] float in [0, 255]. Spatial resizing is applied if
dimensions differ. Returns the aligned gen_depth clipped to [0, 255].
"""
T = min(gen_depth.shape[0], cond_depth.shape[0])
gen_depth = gen_depth[:T]
cond_depth = cond_depth[:T]
H_c, W_c = cond_depth.shape[1], cond_depth.shape[2]
H_g, W_g = gen_depth.shape[1], gen_depth.shape[2]
if (H_g, W_g) != (H_c, W_c):
gen_t = torch.from_numpy(gen_depth).unsqueeze(1).float()
gen_t = F.interpolate(gen_t, size=(H_c, W_c), mode="bilinear", align_corners=False)
gen_depth = gen_t.squeeze(1).numpy()
if valid_mask is None:
valid_mask = cond_depth > 0
else:
valid_mask = valid_mask[:T].astype(bool) & (cond_depth > 0)
if not valid_mask.any():
raise ValueError("No valid pixels remain for depth scale/shift alignment.")
gen_valid = gen_depth[valid_mask].reshape(-1, 1).astype(np.float64)
cond_valid = cond_depth[valid_mask].reshape(-1).astype(np.float64)
A = np.concatenate([gen_valid, np.ones_like(gen_valid)], axis=1)
x, _, _, _ = np.linalg.lstsq(A, cond_valid, rcond=None)
scale, shift = float(x[0]), float(x[1])
return np.clip(scale * gen_depth + shift, 0.0, 255.0).astype(np.float32)
def compute_metrics(
gen_depth: np.ndarray,
cond_depth: np.ndarray,
exclude_mask: np.ndarray | None = None,
) -> tuple[list[float], np.ndarray]:
"""Compute depth metrics between generated and condition depth in grayscale space.
Both inputs are in [0, 255] grayscale. A least-squares scale+shift alignment is
applied to gen_depth before computing metrics, compensating for the affine ambiguity
introduced by per-video depth normalization. Pixels where the condition
depth is zero or ``exclude_mask`` is true are excluded from both alignment
and scoring.
Args:
gen_depth: [T, H, W] float in [0, 255] — depth from generated video
cond_depth: [T, H, W] float in [0, 255] — depth from condition video
exclude_mask: optional [T, H, W] bool mask; true pixels are ignored
Returns:
Tuple of (metric values matching eval_metrics order, aligned gen_depth [T, H, W]).
"""
T = min(gen_depth.shape[0], cond_depth.shape[0])
if exclude_mask is not None:
T = min(T, exclude_mask.shape[0])
cond_depth = cond_depth[:T]
if exclude_mask is not None:
exclude_mask = exclude_mask[:T]
if exclude_mask.shape[1:] != cond_depth.shape[1:]:
mask_t = torch.from_numpy(exclude_mask.astype(np.float32)).unsqueeze(1)
mask_t = F.interpolate(mask_t, size=cond_depth.shape[1:], mode="nearest")
exclude_mask = mask_t.squeeze(1).numpy() > 0.5
valid_mask = cond_depth > 0
if exclude_mask is not None:
valid_mask &= ~exclude_mask.astype(bool)
aligned_gen = _align_depth(gen_depth[:T], cond_depth, valid_mask)
n_valid = valid_mask.sum((-1, -2)) # [T]
valid_frame = n_valid > 0
if not valid_frame.any():
raise ValueError("No valid background pixels remain for depth evaluation.")
pred_ts = torch.from_numpy(aligned_gen[valid_frame]).to(device)
gt_ts = torch.from_numpy(cond_depth[valid_frame]).to(device)
valid_ts = torch.from_numpy(valid_mask[valid_frame]).to(device)
metric_funcs = [getattr(metric, m) for m in eval_metrics]
return [fn(pred_ts, gt_ts, valid_ts).item() for fn in metric_funcs], aligned_gen
# ---------------------------------------------------------------------------
# Side-by-side visualization
# ---------------------------------------------------------------------------
def _save_side_by_side_video(
gen_video: torch.Tensor,
cond_depth: np.ndarray,
aligned_depth: np.ndarray,
output_path: Path,
fps: float,
exclude_mask: np.ndarray | None = None,
) -> None:
"""Write a 3-panel video: generated RGB | condition depth | aligned extracted depth.
Args:
gen_video: [T, C, H, W] float in [0, 1]
cond_depth: [T, H, W] float in [0, 255]
aligned_depth: [T, H, W] float in [0, 255]
output_path: destination .mp4 path
fps: frame rate for the output video
"""
T = min(gen_video.shape[0], cond_depth.shape[0], aligned_depth.shape[0])
gen_np = (gen_video[:T].permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8)
cond_np = np.stack([cond_depth[:T]] * 3, axis=-1).clip(0, 255).astype(np.uint8)
align_np = np.stack([aligned_depth[:T]] * 3, axis=-1).clip(0, 255).astype(np.uint8)
if exclude_mask is not None:
mask = exclude_mask[:T].astype(bool)
# Make mask quality easy to inspect: red overlay on RGB, black on depth.
gen_np[mask] = (
0.45 * gen_np[mask] + 0.55 * np.array([255, 0, 0])
).astype(np.uint8)
cond_np[mask] = 0
align_np[mask] = 0
# Resize panels to the same height if needed (cond/aligned should already match gen)
H = gen_np.shape[1]
def _resize_panel(panel: np.ndarray, height: int) -> np.ndarray:
if panel.shape[1] == height:
return panel
t = torch.from_numpy(panel).permute(0, 3, 1, 2).float()
scale = height / panel.shape[1]
new_w = round(panel.shape[2] * scale)
t = F.interpolate(t, size=(height, new_w), mode="bilinear", align_corners=False)
return t.permute(0, 2, 3, 1).byte().numpy()
cond_np = _resize_panel(cond_np, H)
align_np = _resize_panel(align_np, H)
combined = np.concatenate([gen_np, cond_np, align_np], axis=2) # [T, H, W*3, 3]
output_path.parent.mkdir(parents=True, exist_ok=True)
with av.open(str(output_path), mode="w") as container:
stream = container.add_stream("libx264", rate=Fraction(fps).limit_denominator(1001))
stream.width = combined.shape[2]
stream.height = combined.shape[1]
stream.pix_fmt = "yuv420p"
for frame_np in combined:
frame = av.VideoFrame.from_ndarray(frame_np, format="rgb24")
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
print(f" -> {output_path}")
# ---------------------------------------------------------------------------
# Main evaluation loop
# ---------------------------------------------------------------------------
def evaluate_depth(
generated_dir: Path,
dataset_dir: Path,
depth_repo_path: Path,
encoder: str = "vitl",
input_size: int = 518,
output_dir: Path | None = None,
show_video: bool = False,
exclude_mask_dir: Path | None = None,
shuffle_refs: bool = False,
) -> None:
print(f"Loading depth model ({encoder}) on {device}...")
model = _load_depth_model(encoder, device, depth_repo_path)
cond_dir = dataset_dir
if not cond_dir.exists():
raise FileNotFoundError(f"Condition depth folder not found: {cond_dir}")
generated_files = {f.stem: f for f in sorted(generated_dir.glob("*.mp4"))}
cond_files = {f.stem: f for f in sorted(cond_dir.glob("*.mp4"))}
sample_ids = sorted(set(generated_files) & set(cond_files))
if not sample_ids:
print(f"No matching samples between {generated_dir} and {cond_dir}")
return
# Video-Depth-Anything on the generated clip dominates the cost and does not depend on
# which condition it is compared against, so the shuffled chance level only costs one
# extra condition decode plus a second scale/shift fit and score.
pairing = deranged_pairing(sample_ids) if shuffle_refs else {}
if shuffle_refs and not pairing:
print(f" Only {len(sample_ids)} sample(s); no derangement exists, skipping chance level.")
metric_names = list(eval_metrics) + ([f"{m}_chance" for m in eval_metrics] if pairing else [])
print(f"Evaluating {len(sample_ids)} sample(s)" + (" (+ shuffled chance)" if pairing else "") + "...")
per_sample: dict[str, dict] = {}
results_all: list[list[float]] = []
masking_stats: dict[str, dict] = {}
for i, sid in enumerate(tqdm(sample_ids)):
print(f" [{i + 1}/{len(sample_ids)}] {sid}", end="", flush=True)
gen_video, gen_fps = _read_video(generated_files[sid])
_, _, gen_h, gen_w = gen_video.shape
gen_depth = _extract_depth(gen_video, gen_fps, model, device, input_size)
cond_depth, _ = _load_condition_depth(cond_files[sid], gen_h, gen_w)
exclude_mask = None
if exclude_mask_dir is not None:
mask_path = exclude_mask_dir / f"{sid}.npz"
if not mask_path.exists():
raise FileNotFoundError(
f"Person exclusion mask missing for {sid}: {mask_path}"
)
with np.load(mask_path) as mask_data:
exclude_mask = mask_data["exclude_mask"].astype(bool)
masking_stats[sid] = {
"excluded_fraction": float(exclude_mask.mean()),
"mask_frames": int(exclude_mask.shape[0]),
}
sample_metrics, aligned_depth = compute_metrics(
gen_depth, cond_depth, exclude_mask
)
if pairing:
# The same person-exclusion mask is kept: it describes the GENERATED video, so
# it stays valid when only the reference changes. A shuffled depth can leave no
# valid pixels (its cond_depth > 0 region may not overlap); that sample simply
# has no chance level rather than aborting the run.
other_depth, _ = _load_condition_depth(cond_files[pairing[sid]], gen_h, gen_w)
try:
chance_metrics, _ = compute_metrics(gen_depth, other_depth, exclude_mask)
except ValueError:
chance_metrics = [float("nan")] * len(eval_metrics)
sample_metrics = sample_metrics + chance_metrics
results_all.append(sample_metrics)
per_sample[sid] = dict(zip(metric_names, sample_metrics))
if show_video:
vis_dir = output_dir / "video" if output_dir else generated_dir / "video"
_save_side_by_side_video(
gen_video, cond_depth, aligned_depth,
vis_dir / f"{sid}.mp4", gen_fps,
exclude_mask,
)
metric_str = ", ".join(f"{n}={v:.4f}" for n, v in zip(metric_names, sample_metrics))
print(f" {metric_str}")
final_results = np.array(results_all)
final_results_mean = np.nanmean(final_results, axis=0)
print(f"\n{'=' * 50}")
print(f"Mean metrics over {len(sample_ids)} sample(s):")
for mname, mval in zip(metric_names, final_results_mean):
print(f" {mname}: {mval:.6f}")
if pairing:
means = dict(zip(metric_names, final_results_mean.tolist()))
for m in eval_metrics: # both depth metrics are errors: lower is better
margin = means[f"{m}_chance"] - means[m]
print(f" {m}: chance - signal = {margin:+.6f} [{'OK' if margin > 0 else 'DEGENERATE'}]")
if output_dir is not None:
output_dir.mkdir(parents=True, exist_ok=True)
json_path = output_dir / "results.json"
result = {
"mean": dict(zip(metric_names, final_results_mean.tolist())),
"per_sample": per_sample,
}
if exclude_mask_dir is not None:
result["masking"] = {
"exclude_mask_dir": str(exclude_mask_dir),
"per_sample": masking_stats,
}
with open(json_path, "w") as f:
json.dump(result, f, indent=2)
csv_path = output_dir / "results.csv"
with open(csv_path, "w", newline="") as f:
writer = csv.writer(f)
writer.writerow(["sample_id"] + metric_names)
writer.writerow(["average"] + final_results_mean.tolist())
for sid in sample_ids:
writer.writerow([sid] + [per_sample[sid][m] for m in metric_names])
print(f"Results saved to {output_dir}")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main() -> None:
parser = argparse.ArgumentParser(
description="Evaluate depth consistency of generated videos against condition depth videos.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
"--generated-dir",
required=True,
help="Directory containing generated .mp4 files (one per sample ID).",
)
parser.add_argument(
"--dataset-dir",
required=True,
help=(
"FlexBench dataset directory. Expects a 'depth/' sub-folder with condition "
"depth videos whose stems match the generated file stems."
),
)
parser.add_argument(
"--depth-repo-path",
required=True,
help="Path to a local clone of Video-Depth-Anything.",
)
parser.add_argument(
"--encoder",
default="vitl",
choices=["vits", "vitb", "vitl"],
help="Encoder size for the depth model.",
)
parser.add_argument(
"--input-size",
type=int,
default=518,
help="Input resolution for depth model inference.",
)
parser.add_argument(
"--output-dir",
default=None,
help="Optional folder to write results.json and results.csv.",
)
parser.add_argument(
"--exclude-mask-dir",
default=None,
help=(
"Optional directory of <id>.npz masks from generate_person_masks.py. "
"Masked pixels are excluded from scale/shift alignment and metrics."
),
)
parser.add_argument(
"--shuffle-refs",
action="store_true",
help=(
"Also score every video against a different sample's condition depth and "
"report it as <metric>_chance. Reuses the Video-Depth-Anything extraction."
),
)
parser.add_argument(
"--show-video",
action="store_true",
help=(
"Save a side-by-side comparison video for each sample: "
"generated RGB | condition depth | aligned extracted depth. "
"Written to <output-dir>/vis/<id>.mp4 (or <generated-dir>/vis/<id>.mp4)."
),
)
args = parser.parse_args()
evaluate_depth(
generated_dir=Path(args.generated_dir),
dataset_dir=Path(args.dataset_dir),
depth_repo_path=Path(args.depth_repo_path),
encoder=args.encoder,
input_size=args.input_size,
output_dir=Path(args.output_dir) if args.output_dir else None,
show_video=args.show_video,
exclude_mask_dir=Path(args.exclude_mask_dir) if args.exclude_mask_dir else None,
shuffle_refs=args.shuffle_refs,
)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
21.6 kB
·
Xet hash:
2500f778dc87890bb4bd1e1f84b4b777f317922ef779abc03b79a6903b473d48

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.