import argparse import time from pathlib import Path import cv2 import numpy as np import onnxruntime as ort from PIL import Image IMG_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".webp"} def get_providers(force_cpu: bool = False): available = ort.get_available_providers() if force_cpu: return ["CPUExecutionProvider"] providers = [] if "CUDAExecutionProvider" in available: providers.append("CUDAExecutionProvider") providers.append("CPUExecutionProvider") return providers def print_model_io(session: ort.InferenceSession): print("\n[MODEL INPUTS]") for i, inp in enumerate(session.get_inputs()): print(f" {i}: name={inp.name} shape={inp.shape} type={inp.type}") print("\n[MODEL OUTPUTS]") for i, out in enumerate(session.get_outputs()): print(f" {i}: name={out.name} shape={out.shape} type={out.type}") print("") def resolve_hw_from_input_shape(shape, fallback_h: int, fallback_w: int): """ Tenta descobrir H/W do input ONNX. Esperado normalmente: [1, 3, H, W] ou [1, H, W, 3]. Se for dinâmico, usa fallback. """ if shape is None: return fallback_h, fallback_w dims = list(shape) def is_int(x): return isinstance(x, int) and x > 0 # NCHW if len(dims) == 4 and dims[1] == 3: h = dims[2] if is_int(dims[2]) else fallback_h w = dims[3] if is_int(dims[3]) else fallback_w return int(h), int(w) # NHWC if len(dims) == 4 and dims[-1] == 3: h = dims[1] if is_int(dims[1]) else fallback_h w = dims[2] if is_int(dims[2]) else fallback_w return int(h), int(w) return fallback_h, fallback_w def is_nchw_input(shape): if shape is None: return True dims = list(shape) return len(dims) == 4 and dims[1] == 3 def preprocess_image(image_path: Path, input_shape, input_h: int, input_w: int, mean, std): img_pil = Image.open(image_path).convert("RGB") rgb_orig = np.array(img_pil) bgr_orig = cv2.cvtColor(rgb_orig, cv2.COLOR_RGB2BGR) rgb_resized = cv2.resize(rgb_orig, (input_w, input_h), interpolation=cv2.INTER_AREA) x = rgb_resized.astype(np.float32) / 255.0 mean_arr = np.array(mean, dtype=np.float32).reshape(1, 1, 3) std_arr = np.array(std, dtype=np.float32).reshape(1, 1, 3) x = (x - mean_arr) / std_arr if is_nchw_input(input_shape): x = np.transpose(x, (2, 0, 1)) # CHW x = np.expand_dims(x, axis=0).astype(np.float32) return bgr_orig, x def robust_normalize_for_view(depth_m: np.ndarray, invert: bool = False): """ Só para visualização colorida. O modo bandas usa o depth_m direto. """ d = depth_m.astype(np.float32) valid = np.isfinite(d) & (d > 0) if np.count_nonzero(valid) < 20: return np.zeros_like(d, dtype=np.float32) vals = d[valid] p2 = np.percentile(vals, 2) p98 = np.percentile(vals, 98) dn = (d - p2) / (p98 - p2 + 1e-6) dn = np.clip(dn, 0.0, 1.0) if invert: dn = 1.0 - dn dn[~valid] = 0.0 return dn.astype(np.float32) def depth_to_colormap(depth_norm: np.ndarray): u8 = (depth_norm * 255).astype(np.uint8) return cv2.applyColorMap(u8, cv2.COLORMAP_TURBO) def make_metric_bands(depth_m: np.ndarray, near_m: float, mid_m: float, far_m: float): """ Bandas baseadas no valor métrico estimado. 0: < near_m 1: near_m até mid_m 2: mid_m até far_m 3: >= far_m """ d = depth_m.astype(np.float32) valid = np.isfinite(d) & (d > 0) bands = np.zeros_like(d, dtype=np.uint8) bands[(d >= near_m) & (d < mid_m)] = 1 bands[(d >= mid_m) & (d < far_m)] = 2 bands[d >= far_m] = 3 bands[~valid] = 255 out = np.zeros((bands.shape[0], bands.shape[1], 3), dtype=np.uint8) # BGR out[bands == 0] = (40, 40, 255) # muito perto out[bands == 1] = (40, 180, 255) # perto/médio out[bands == 2] = (40, 255, 120) # médio/longe out[bands == 3] = (255, 180, 40) # longe out[bands == 255] = (0, 0, 0) # inválido return out def extract_depth_from_outputs(outputs, output_names): """ Tenta achar depth em diferentes formatos: - output chamado depth - output chamado points/xyz, usando canal Z - saída única """ name_to_out = { name.lower(): out for name, out in zip(output_names, outputs) } # 1) Procura saída com nome depth for name, out in name_to_out.items(): if "depth" in name: return squeeze_depth(out) # 2) Procura points/xyz e usa Z for name, out in name_to_out.items(): if "point" in name or "xyz" in name: arr = np.array(out) return extract_z_from_points(arr) # 3) Se tiver só uma saída, usa ela if len(outputs) == 1: arr = np.array(outputs[0]) # Se parece points, extrai Z if arr.ndim == 4 and (arr.shape[1] == 3 or arr.shape[-1] == 3): return extract_z_from_points(arr) return squeeze_depth(arr) # 4) Fallback: pega a primeira saída 2D/3D/4D plausível for out in outputs: arr = np.array(out) try: d = squeeze_depth(arr) if d.ndim == 2: return d except Exception: pass raise RuntimeError("Não consegui identificar o mapa de depth nas saídas ONNX.") def squeeze_depth(arr): arr = np.array(arr) # Remove batch/canal unitário arr = np.squeeze(arr) if arr.ndim == 2: return arr.astype(np.float32) if arr.ndim == 3: # CHW com 1 canal if arr.shape[0] == 1: return arr[0].astype(np.float32) # HWC com 1 canal if arr.shape[-1] == 1: return arr[..., 0].astype(np.float32) # Se for 3 canais, talvez seja XYZ if arr.shape[0] == 3: return arr[2].astype(np.float32) if arr.shape[-1] == 3: return arr[..., 2].astype(np.float32) raise RuntimeError(f"Formato de depth não suportado: shape={arr.shape}") def extract_z_from_points(arr): arr = np.array(arr) # NCHW: [1, 3, H, W] if arr.ndim == 4 and arr.shape[1] == 3: return arr[0, 2].astype(np.float32) # NHWC: [1, H, W, 3] if arr.ndim == 4 and arr.shape[-1] == 3: return arr[0, ..., 2].astype(np.float32) # CHW: [3, H, W] if arr.ndim == 3 and arr.shape[0] == 3: return arr[2].astype(np.float32) # HWC: [H, W, 3] if arr.ndim == 3 and arr.shape[-1] == 3: return arr[..., 2].astype(np.float32) raise RuntimeError(f"Formato de points/xyz não suportado: shape={arr.shape}") def maybe_invert_depth_if_needed(depth: np.ndarray, auto_invert: bool): """ Alguns modelos podem devolver inverso/disparity. Aqui deixei opcional. Por padrão não mexe. """ if not auto_invert: return depth d = depth.astype(np.float32) valid = np.isfinite(d) & (d > 1e-6) out = np.zeros_like(d, dtype=np.float32) out[valid] = 1.0 / d[valid] return out def infer_depth(session, image_path: Path, args): input0 = session.get_inputs()[0] input_shape = input0.shape fallback_h, fallback_w = args.input_h, args.input_w input_h, input_w = resolve_hw_from_input_shape(input_shape, fallback_h, fallback_w) preview_bgr, x = preprocess_image( image_path=image_path, input_shape=input_shape, input_h=input_h, input_w=input_w, mean=args.mean, std=args.std, ) feed = {} # Alimenta input principal feed[input0.name] = x # Alguns modelos têm input extra de intrinsics/K. # Se existir, manda uma matriz aproximada baseada no tamanho de entrada. # Para nosso teste visual, isso é melhor do que travar. for inp in session.get_inputs()[1:]: name = inp.name.lower() shape = inp.shape fx = args.fx if args.fx > 0 else input_w * 0.9 fy = args.fy if args.fy > 0 else input_w * 0.9 cx = args.cx if args.cx >= 0 else input_w / 2.0 cy = args.cy if args.cy >= 0 else input_h / 2.0 K = np.array( [ [fx, 0.0, cx], [0.0, fy, cy], [0.0, 0.0, 1.0], ], dtype=np.float32, ) if "k" == name or "intr" in name or "camera" in name: if len(shape) == 3: feed[inp.name] = K[None, ...] else: feed[inp.name] = K else: # Fallback para input extra desconhecido # Evita crash, mas imprime para a gente ajustar se precisar. print(f"[WARN] Input extra desconhecido: {inp.name}, shape={inp.shape}. Enviando zeros.") concrete_shape = [] for dim in shape: concrete_shape.append(dim if isinstance(dim, int) and dim > 0 else 1) feed[inp.name] = np.zeros(concrete_shape, dtype=np.float32) t0 = time.time() outputs = session.run(None, feed) infer_ms = (time.time() - t0) * 1000.0 output_names = [o.name for o in session.get_outputs()] depth = extract_depth_from_outputs(outputs, output_names) depth = maybe_invert_depth_if_needed(depth, auto_invert=args.inv_depth) # Redimensiona para o preview original if depth.shape[:2] != preview_bgr.shape[:2]: depth = cv2.resize( depth.astype(np.float32), (preview_bgr.shape[1], preview_bgr.shape[0]), interpolation=cv2.INTER_CUBIC, ) # Remove valores absurdos só para visual/estatística depth = depth.astype(np.float32) depth[~np.isfinite(depth)] = 0.0 depth[depth < 0] = 0.0 return preview_bgr, depth, infer_ms def calc_stats(depth_m: np.ndarray, infer_ms: float): valid = np.isfinite(depth_m) & (depth_m > 0) if np.count_nonzero(valid) < 20: return { "p01": 0.0, "p05": 0.0, "p50": 0.0, "p95": 0.0, "p99": 0.0, "mean": 0.0, "std": 0.0, "infer_ms": infer_ms, "valid_pct": 0.0, } vals = depth_m[valid] return { "p01": float(np.percentile(vals, 1)), "p05": float(np.percentile(vals, 5)), "p50": float(np.percentile(vals, 50)), "p95": float(np.percentile(vals, 95)), "p99": float(np.percentile(vals, 99)), "mean": float(np.mean(vals)), "std": float(np.std(vals)), "infer_ms": float(infer_ms), "valid_pct": float(np.mean(valid) * 100.0), } def resize_to_height(img: np.ndarray, target_h: int): h, w = img.shape[:2] if h == target_h: return img scale = target_h / h new_w = max(1, int(w * scale)) return cv2.resize(img, (new_w, target_h), interpolation=cv2.INTER_AREA) def draw_header(canvas, lines): header_h = 24 + 24 * len(lines) cv2.rectangle(canvas, (0, 0), (canvas.shape[1], header_h), (0, 0, 0), -1) y = 24 for line in lines: cv2.putText( canvas, line, (12, y), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (255, 255, 255), 1, cv2.LINE_AA, ) y += 24 return canvas def compose_view(preview_bgr, depth_m, image_path, idx, total, mode, invert_view, stats, args): depth_norm = robust_normalize_for_view(depth_m, invert=invert_view) depth_color = depth_to_colormap(depth_norm) bands_color = make_metric_bands(depth_m, args.near_m, args.mid_m, args.far_m) right_src = depth_color if mode == "depth" else bands_color left = resize_to_height(preview_bgr, args.view_h) right = resize_to_height(right_src, args.view_h) h = min(left.shape[0], right.shape[0]) left = left[:h] right = right[:h] canvas = np.hstack([left, right]) mode_name = "depth colormap" if mode == "depth" else "metric bands" lines = [ f"{idx + 1}/{total} - {image_path.name}", f"modo={mode_name} | p05={stats['p05']:.2f}m p50={stats['p50']:.2f}m p95={stats['p95']:.2f}m valid={stats['valid_pct']:.1f}% infer={stats['infer_ms']:.1f}ms", f"bands: <{args.near_m:.2f}m | {args.near_m:.2f}-{args.mid_m:.2f}m | {args.mid_m:.2f}-{args.far_m:.2f}m | >{args.far_m:.2f}m", "N/SPACE prox | A ant | M modo | I inverte visual | S salvar | Q sair", ] return draw_header(canvas, lines) def parse_mean_std(text): vals = [float(x.strip()) for x in text.split(",")] if len(vals) != 3: raise argparse.ArgumentTypeError("Use 3 valores separados por vírgula. Ex: 0.485,0.456,0.406") return vals def main(): parser = argparse.ArgumentParser() parser.add_argument("--input_dir", required=True, help="Pasta com previews PNG/JPG") parser.add_argument("--model_path", required=True, help="Caminho do arquivo .onnx") parser.add_argument("--cpu", action="store_true", help="Força CPUExecutionProvider") # Usado se o modelo tiver tamanho dinâmico parser.add_argument("--input_h", type=int, default=384) parser.add_argument("--input_w", type=int, default=512) # Normalização padrão ImageNet. Se o repo do ONNX pedir outra, ajustamos. parser.add_argument("--mean", type=parse_mean_std, default=[0.485, 0.456, 0.406]) parser.add_argument("--std", type=parse_mean_std, default=[0.229, 0.224, 0.225]) # Intrínsecos aproximados se o ONNX pedir K/intrinsics parser.add_argument("--fx", type=float, default=-1.0) parser.add_argument("--fy", type=float, default=-1.0) parser.add_argument("--cx", type=float, default=-1.0) parser.add_argument("--cy", type=float, default=-1.0) # Se o output for inverso/disparity, ativa isto parser.add_argument("--inv_depth", action="store_true") # Bandas métricas para visual parser.add_argument("--near_m", type=float, default=0.6) parser.add_argument("--mid_m", type=float, default=1.1) parser.add_argument("--far_m", type=float, default=1.8) parser.add_argument("--view_h", type=int, default=560) parser.add_argument("--save_dir", default="unidepth_onnx_viewer_saves") args = parser.parse_args() input_dir = Path(args.input_dir) model_path = Path(args.model_path) save_dir = Path(args.save_dir) save_dir.mkdir(parents=True, exist_ok=True) if not model_path.exists(): raise FileNotFoundError(f"Modelo ONNX não encontrado: {model_path}") image_paths = sorted([ p for p in input_dir.rglob("*") if p.suffix.lower() in IMG_EXTS ]) if not image_paths: raise RuntimeError(f"Nenhuma imagem encontrada em: {input_dir}") providers = get_providers(force_cpu=args.cpu) print(f"[INFO] imagens: {len(image_paths)}") print(f"[INFO] model_path: {model_path}") print(f"[INFO] providers: {providers}") print(f"[INFO] ONNX Runtime providers disponiveis: {ort.get_available_providers()}") session = ort.InferenceSession(str(model_path), providers=providers) print_model_io(session) idx = 0 mode = "depth" invert_view = False cached_path = None cached_data = None cv2.namedWindow("UniDepth ONNX Metric Viewer", cv2.WINDOW_NORMAL) while True: image_path = image_paths[idx] if cached_path != image_path or cached_data is None: print(f"[RUN] {idx + 1}/{len(image_paths)} - {image_path.name}") try: preview_bgr, depth_m, infer_ms = infer_depth(session, image_path, args) stats = calc_stats(depth_m, infer_ms) cached_data = (preview_bgr, depth_m, stats) cached_path = image_path except Exception as e: print(f"[ERRO] Falha em {image_path}: {e}") idx = min(idx + 1, len(image_paths) - 1) cached_path = None cached_data = None continue else: preview_bgr, depth_m, stats = cached_data view = compose_view( preview_bgr=preview_bgr, depth_m=depth_m, image_path=image_path, idx=idx, total=len(image_paths), mode=mode, invert_view=invert_view, stats=stats, args=args, ) cv2.imshow("UniDepth ONNX Metric Viewer", view) key = cv2.waitKey(0) & 0xFF if key in [27, ord("q"), ord("Q")]: break elif key in [ord("n"), ord("N"), 32]: idx = min(idx + 1, len(image_paths) - 1) cached_path = None elif key in [ord("a"), ord("A")]: idx = max(idx - 1, 0) cached_path = None elif key in [ord("m"), ord("M")]: mode = "bands" if mode == "depth" else "depth" elif key in [ord("i"), ord("I")]: invert_view = not invert_view print(f"[INFO] invert_view={invert_view}") elif key in [ord("s"), ord("S")]: out_path = save_dir / f"{image_path.stem}_unidepth_onnx_{mode}.png" cv2.imwrite(str(out_path), view) print(f"[SAVE] {out_path}") cv2.destroyAllWindows() if __name__ == "__main__": main()