back / Qwen3-VL /sim_scripts /metrics.py
shenaosdfa's picture
Upload Qwen3-VL code only
912284f verified
Raw
History Blame Contribute Delete
4.63 kB
"""Mesh distance metrics for garment reward computation.
Provides chamfer distance between meshes with potentially different topology
(different vertex/face counts), suitable for comparing a predicted simulated
mesh against a ground-truth mesh from GarmentCodeData.
"""
from __future__ import annotations
from pathlib import Path
from typing import Optional
import numpy as np
def load_mesh_vertices(path: str | Path) -> np.ndarray:
"""Load vertex positions from OBJ or PLY files.
Returns (N, 3) float64 array of vertex positions.
"""
path = Path(path)
suffix = path.suffix.lower()
if suffix == ".obj":
return _load_obj_vertices(path)
elif suffix == ".ply":
return _load_ply_vertices(path)
else:
raise ValueError(f"Unsupported mesh format: {suffix} ({path})")
def _load_obj_vertices(path: Path) -> np.ndarray:
verts = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
if line.startswith("v "):
parts = line.strip().split()
verts.append([float(parts[1]), float(parts[2]), float(parts[3])])
if not verts:
raise ValueError(f"No vertices found in {path}")
return np.array(verts, dtype=np.float64)
def _load_ply_vertices(path: Path) -> np.ndarray:
verts = []
in_header = True
vertex_count = 0
with open(path, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if in_header:
if line.startswith("element vertex"):
vertex_count = int(line.split()[-1])
elif line == "end_header":
in_header = False
continue
if vertex_count <= 0:
break
parts = line.split()
if len(parts) >= 3:
verts.append([float(parts[0]), float(parts[1]), float(parts[2])])
vertex_count -= 1
if not verts:
raise ValueError(f"No vertices found in {path}")
return np.array(verts, dtype=np.float64)
def chamfer_distance(
pred_verts: np.ndarray,
gt_verts: np.ndarray,
batch_size: int = 4096,
) -> float:
"""Symmetric chamfer distance between two point clouds.
Handles different-topology meshes (different vertex counts).
Uses batched computation to avoid OOM on large meshes.
Args:
pred_verts: (N, 3) predicted vertex positions
gt_verts: (M, 3) ground truth vertex positions
batch_size: process this many query points at a time
Returns:
Scalar symmetric chamfer distance (mean of both directions).
"""
pred = np.asarray(pred_verts, dtype=np.float64)
gt = np.asarray(gt_verts, dtype=np.float64)
pred_to_gt = _directed_chamfer(pred, gt, batch_size)
gt_to_pred = _directed_chamfer(gt, pred, batch_size)
return float(pred_to_gt + gt_to_pred)
def _directed_chamfer(
source: np.ndarray, target: np.ndarray, batch_size: int
) -> float:
"""Mean of min distances from each source point to the nearest target point."""
n = source.shape[0]
min_dists = np.empty(n, dtype=np.float64)
for start in range(0, n, batch_size):
end = min(start + batch_size, n)
chunk = source[start:end] # (B, 3)
diff = chunk[:, None, :] - target[None, :, :] # (B, M, 3)
dist_sq = (diff * diff).sum(axis=-1) # (B, M)
min_dists[start:end] = dist_sq.min(axis=1)
return float(np.sqrt(min_dists).mean())
def chamfer_distance_from_files(
pred_path: str | Path,
gt_path: str | Path,
) -> float:
"""Convenience: load meshes from files and compute chamfer distance."""
pred_verts = load_mesh_vertices(pred_path)
gt_verts = load_mesh_vertices(gt_path)
return chamfer_distance(pred_verts, gt_verts)
def find_gt_sim_mesh(
sample_id: str,
gcd_root: str | Path,
) -> Optional[Path]:
"""Find the ground-truth simulated mesh for a GarmentCodeData sample.
Looks for the pattern: <gcd_root>/**/default_body/<sample_id>/<sample_id>_sim.ply
Falls back to _sim.obj if .ply not found.
"""
gcd_root = Path(gcd_root)
for garments_dir in sorted(gcd_root.iterdir()):
if not garments_dir.is_dir():
continue
body_dir = garments_dir / "default_body"
if not body_dir.is_dir():
continue
sample_dir = body_dir / sample_id
if not sample_dir.is_dir():
continue
for suffix in (".ply", ".obj"):
mesh = sample_dir / f"{sample_id}_sim{suffix}"
if mesh.exists():
return mesh
return None