fishxinyu's picture
download
raw
11.3 kB
"""Evaluate generated videos against condition inputs.
For --type depth:
- Loads the condition depth video from <dataset-dir>/depth/<id>.mp4
- Applies the same resize-to-fill + center-crop that ic_lora.py uses when
feeding the condition to the model (matching the generated video's resolution)
- Extracts depth from the generated video using Video-Depth-Anything
- Both depths are normalized per-video to [0, 1] (same pipeline used to create
the dataset conditions via compute_reference.py::compute_depth_reference)
- Computes per-sample MSE and reports the mean
Usage:
python evaluation/evaluate.py \
--type depth \
--dataset-dir /home/xinyuy/dataset_processing/flexcombine-bench/datasets/v0/subject_depth \
--generated-dir /home/xinyuy/dataset_processing/flexcombine-bench/flexbench/results/v0/subject_depth \
--depth-repo-path tools/Video-Depth-Anything
"""
import argparse
import json
import math
import sys
from pathlib import Path
import av
import numpy as np
import torch
import torch.nn.functional as F
# ---------------------------------------------------------------------------
# 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.
Args:
video: [T, C, H, W] float tensor
height: target height
width: target width
Returns:
[T, C, height, width] float tensor
"""
_, _, 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 helpers
# ---------------------------------------------------------------------------
_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-"
f"{'Small' if encoder == 'vits' else 'Base' if encoder == 'vitb' else 'Large'}"
f" and 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 depth estimation and normalize per-video to [0, 1].
Mirrors the exact preprocessing in compute_reference.py::compute_depth_reference
so that the output is comparable to stored condition depth videos.
Args:
video: [T, C, H, W] float tensor in [0, 1]
Returns:
depths: [T, H, W] float numpy array in [0, 1]
"""
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 # [T, H, W]
def _load_condition_depth(
cond_path: str | Path,
target_height: int,
target_width: int,
) -> tuple[np.ndarray, float]:
"""Load a stored depth condition video, cropped to the generated video's resolution.
Applies the same resize-to-fill + center-crop that ic_lora.py::video_preprocess
uses when feeding the condition to the model, so we compare the same spatial
region that the model actually saw.
The depth condition is stored as a grayscale video (all 3 channels equal to the
normalized depth value). We take channel 0 after reading back as uint8→[0,1].
Returns:
depths: [T, H, W] float array in [0, 1]
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(), fps # [T, H, W]
def _compute_mse(pred: np.ndarray, cond: np.ndarray) -> float:
"""MSE between pred and cond depth arrays (both in [0, 1]).
Trims to the shorter temporal length. Spatial dims should already match
after the center-crop applied in _load_condition_depth.
"""
T = min(pred.shape[0], cond.shape[0])
pred, cond = pred[:T], cond[:T]
H_p, W_p = pred.shape[1], pred.shape[2]
H_c, W_c = cond.shape[1], cond.shape[2]
if (H_p, W_p) != (H_c, W_c):
# Fallback in case resolutions still differ (e.g. odd rounding)
pred_t = torch.from_numpy(pred).unsqueeze(1).float() # [T, 1, H, W]
pred_t = F.interpolate(pred_t, size=(H_c, W_c), mode="bilinear", align_corners=False)
pred = pred_t.squeeze(1).numpy()
return float(np.mean((pred - cond) ** 2))
# ---------------------------------------------------------------------------
# Evaluation entry points
# ---------------------------------------------------------------------------
def evaluate_depth(
generated_dir: Path,
dataset_dir: Path,
depth_repo_path: Path,
encoder: str = "vitl",
input_size: int = 518,
output_json: Path | None = None,
) -> None:
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Loading depth model ({encoder}) on {device}...")
model = _load_depth_model(encoder, device, depth_repo_path)
cond_dir = dataset_dir / "depth"
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
print(f"Evaluating {len(sample_ids)} sample(s)...")
per_sample: dict[str, float] = {}
for i, sid in enumerate(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)
# Apply the same resize-to-fill + center-crop that ic_lora.py applies so we
# compare the depth of the spatial region the model actually conditioned on.
cond_depth, _ = _load_condition_depth(cond_files[sid], gen_h, gen_w)
mse = _compute_mse(gen_depth, cond_depth)
per_sample[sid] = mse
print(f" MSE={mse:.6f}")
mean_mse = float(np.mean(list(per_sample.values())))
print(f"\nMean Depth MSE: {mean_mse:.6f} ({len(per_sample)} samples)")
if output_json is not None:
result = {"mean_mse": mean_mse, "per_sample": per_sample}
output_json.parent.mkdir(parents=True, exist_ok=True)
with open(output_json, "w") as f:
json.dump(result, f, indent=2)
print(f"Results saved to {output_json}")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main() -> None:
parser = argparse.ArgumentParser(
description="Evaluate generated videos against FlexBench condition inputs.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument(
"--type",
required=True,
choices=["depth"],
help="Evaluation type.",
)
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. For --type depth, expects a 'depth/' "
"sub-folder with condition depth videos matching the generated file stems."
),
)
parser.add_argument(
"--depth-repo-path",
default=None,
help="Path to a local clone of Video-Depth-Anything (required for --type depth).",
)
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-json",
default=None,
help="Optional path to write per-sample MSE results as JSON.",
)
args = parser.parse_args()
if args.type == "depth" and args.depth_repo_path is None:
parser.error("--depth-repo-path is required for --type depth")
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_json=Path(args.output_json) if args.output_json else None,
)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
11.3 kB
·
Xet hash:
d2208aaa59f83b747304f0ac414e405a2b6ba9998e27436dd40aa189dabb42cc

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