agrobot_base/Python/OAK/datasets/oak-fcc-3/depth_unimetric_viewer.py

573 lines
17 KiB
Python

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()