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

1315 lines
45 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
_11_validate_onnx.py
Valida fidelidade entre:
- modelo PyTorch .pt
- modelo ONNX .onnx
para o SegFormer OAK-FCC-3 Multi-Head.
Exemplo:
python _11_validate_onnx.py --config config.json --max_samples 20 --device cuda --onnx_provider cuda --torch_no_amp
python .\_11_validate_onnx.py --config .\config.json --train-script .\_8_train_multihead_v2.py --max_samples 50 --device cuda --onnx_provider cuda --torch_no_amp
Para validar o ONNX com saída já redimensionada:
python _11_validate_onnx.py --config config.json --max_samples 20 --device cuda --onnx_provider cuda --torch_no_amp --compare_at_input_size
"""
from __future__ import annotations
import os
import json
import time
import copy
import argparse
import importlib.util
from pathlib import Path
from typing import Dict, List, Optional, Sequence, Tuple
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
SOURCE_CHANNEL_ORDER = ["R", "G", "B", "RE", "NIR"]
DERIVED_CHANNEL_ORDER = ["NDVI", "NDRE"]
SUPPORTED_INPUT_CHANNELS = SOURCE_CHANNEL_ORDER + DERIVED_CHANNEL_ORDER
# ============================================================
# Utils
# ============================================================
def load_json(path: str | Path) -> dict:
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
def save_json(path: str | Path, data: dict):
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
def resolve_path(path_like: Optional[str], base: Optional[Path] = None) -> Optional[Path]:
if path_like is None:
return None
p = Path(path_like)
if p.is_absolute():
return p
if base is None:
base = Path.cwd()
return (base / p).resolve()
def import_train_module(train_script_path: str | Path):
train_script_path = Path(train_script_path)
if not train_script_path.exists():
raise FileNotFoundError(f"Script de treino não encontrado: {train_script_path}")
spec = importlib.util.spec_from_file_location(
"train_multihead_module",
str(train_script_path.resolve())
)
if spec is None or spec.loader is None:
raise RuntimeError(f"Não consegui importar o script: {train_script_path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
def softmax_np(logits: np.ndarray, axis: int = 1) -> np.ndarray:
x = logits.astype(np.float32)
x = x - np.max(x, axis=axis, keepdims=True)
e = np.exp(x)
return e / np.clip(np.sum(e, axis=axis, keepdims=True), 1e-12, None)
def resize_logits_np_nchw(logits: np.ndarray, target_hw: Tuple[int, int]) -> np.ndarray:
"""
logits: [N,C,H,W]
target_hw: (H,W)
"""
n, c, h, w = logits.shape
th, tw = target_hw
if (h, w) == (th, tw):
return logits
out = np.empty((n, c, th, tw), dtype=np.float32)
for bi in range(n):
for ci in range(c):
out[bi, ci] = cv2.resize(
logits[bi, ci].astype(np.float32),
(tw, th),
interpolation=cv2.INTER_LINEAR,
)
return out
def compute_mask_iou_between_preds(
pred_a: np.ndarray,
pred_b: np.ndarray,
num_classes: int,
) -> Tuple[List[Optional[float]], float, List[int]]:
"""
Mede IoU entre duas predições.
Classes ausentes nos dois mapas recebem None e NÃO entram no mIoU.
Isso evita o caso:
máscaras iguais, só classe 0 presente -> [1.0, None, None] -> mIoU=1.0
em vez de:
[1.0, 0.0, 0.0] -> mIoU=0.333
"""
a = pred_a.reshape(-1).astype(np.int64)
b = pred_b.reshape(-1).astype(np.int64)
valid = (a >= 0) & (a < num_classes) & (b >= 0) & (b < num_classes)
a = a[valid]
b = b[valid]
if a.size == 0:
return [None for _ in range(num_classes)], 0.0, []
cm = np.bincount(
num_classes * a + b,
minlength=num_classes * num_classes,
).reshape(num_classes, num_classes)
tp = np.diag(cm).astype(np.float64)
fp = cm.sum(axis=0).astype(np.float64) - tp
fn = cm.sum(axis=1).astype(np.float64) - tp
den = tp + fp + fn
iou_per_class: List[Optional[float]] = []
present_classes: List[int] = []
for cls in range(num_classes):
if den[cls] <= 0:
# Classe ausente nas duas predições.
iou_per_class.append(None)
else:
iou_per_class.append(float(tp[cls] / den[cls]))
present_classes.append(cls)
valid_ious = [x for x in iou_per_class if x is not None]
miou = float(np.mean(valid_ious)) if valid_ious else 0.0
return iou_per_class, miou, present_classes
def load_norm_stats(
path: Optional[Path],
channel_names: Sequence[str],
) -> Tuple[Optional[List[float]], Optional[List[float]], Optional[str]]:
if path is None or not path.is_file():
if path is not None:
print(f"[NORM] norm_stats não encontrado: {path}")
raise FileNotFoundError(
"Validação PyTorch x ONNX exige os mesmos norm_stats do treinamento."
)
js = load_json(path)
mean = js.get("mean")
std = js.get("std")
names = [str(x).strip().upper() for x in js.get("channels", [])]
requested = [str(x).strip().upper() for x in channel_names]
if mean is None or std is None:
raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}")
if len(mean) != len(std):
raise RuntimeError(f"norm_stats inválido: mean={len(mean)} std={len(std)}")
if names:
if len(names) != len(mean):
raise RuntimeError(
f"norm_stats incompatível: channels={len(names)} mean={len(mean)}"
)
if len(set(names)) != len(names):
raise RuntimeError(f"norm_stats possui canais duplicados: {names}")
missing = [name for name in requested if name not in names]
if missing:
raise RuntimeError(
f"norm_stats não contém {missing}. Disponíveis={names}, pedidos={requested}"
)
indices = [names.index(name) for name in requested]
elif len(mean) == len(requested):
indices = list(range(len(requested)))
print(f"[NORM][WARN] stats sem nomes; assumindo ordem {requested}")
elif len(mean) == len(SOURCE_CHANNEL_ORDER) and all(x in SOURCE_CHANNEL_ORDER for x in requested):
indices = [SOURCE_CHANNEL_ORDER.index(name) for name in requested]
print(f"[NORM][WARN] stats Raw5 legados; usando ordem {SOURCE_CHANNEL_ORDER}")
else:
raise RuntimeError(
f"Não é seguro mapear norm_stats sem nomes: mean={len(mean)}, pedidos={requested}"
)
mean_sel = [float(mean[i]) for i in indices]
std_sel = [float(std[i]) for i in indices]
if not np.all(np.isfinite(mean_sel)) or not np.all(np.isfinite(std_sel)):
raise RuntimeError(f"norm_stats possui NaN/Inf: mean={mean_sel} std={std_sel}")
if any(x <= 0.0 for x in std_sel):
raise RuntimeError(f"norm_stats possui std inválido: {std_sel}")
print(f"[NORM] usando {path}")
print(f"[NORM] channels={requested}")
print(f"[NORM] mean={mean_sel}")
print(f"[NORM] std ={std_sel}")
return mean_sel, std_sel, str(path)
def normalize_numpy_chw(chw: np.ndarray, mean: Optional[List[float]], std: Optional[List[float]]) -> np.ndarray:
if mean is None or std is None:
return chw.astype(np.float32)
mean_np = np.asarray(mean, dtype=np.float32).reshape(-1, 1, 1)
std_np = np.asarray(std, dtype=np.float32).reshape(-1, 1, 1)
std_np = np.clip(std_np, 1e-6, None)
return ((chw.astype(np.float32) - mean_np) / std_np).astype(np.float32)
def load_tensor(
path: Path,
requested_channels: Sequence[str],
saved_channel_names: Optional[Sequence[str]] = None,
) -> np.ndarray:
arr = np.load(str(path)).astype(np.float32)
if arr.ndim != 3:
raise RuntimeError(f"Tensor inválido {path}: shape={arr.shape}, esperado 3D")
if 1 <= arr.shape[0] <= len(SUPPORTED_INPUT_CHANNELS) and arr.shape[1] > 8 and arr.shape[2] > 8:
chw = arr
elif 1 <= arr.shape[-1] <= len(SUPPORTED_INPUT_CHANNELS) and arr.shape[0] > 8 and arr.shape[1] > 8:
chw = np.transpose(arr, (2, 0, 1))
else:
raise RuntimeError(f"Tensor com layout inesperado: {path} shape={arr.shape}")
requested = [str(x).strip().upper() for x in requested_channels]
saved = [str(x).strip().upper() for x in (saved_channel_names or [])]
if saved:
if len(saved) != chw.shape[0] or len(set(saved)) != len(saved):
raise RuntimeError(
f"Meta/tensor incompatíveis em {path}: channels={saved}, shape={chw.shape}"
)
missing = [name for name in requested if name not in saved]
if missing:
raise RuntimeError(f"Tensor {path} não contém {missing}. Disponíveis={saved}")
indices = [saved.index(name) for name in requested]
elif chw.shape[0] == len(requested):
indices = list(range(len(requested)))
print(f"[WARN] {path.name}: sem meta.channels; assumindo ordem {requested}")
elif chw.shape[0] == len(SOURCE_CHANNEL_ORDER) and all(x in SOURCE_CHANNEL_ORDER for x in requested):
indices = [SOURCE_CHANNEL_ORDER.index(name) for name in requested]
print(f"[WARN] {path.name}: Raw5 legado; usando ordem {SOURCE_CHANNEL_ORDER}")
else:
raise RuntimeError(
f"Não é seguro inferir canais de {path}: shape={chw.shape}, pedidos={requested}"
)
chw = chw[indices, :, :]
finite = np.isfinite(chw)
if finite.any():
mx = float(np.nanmax(chw[finite]))
if mx > 2.0 and mx <= 255.0:
chw = chw / 255.0
elif mx > 255.0:
chw = chw / 65535.0
chw = np.nan_to_num(chw, nan=0.0, posinf=1.0, neginf=-1.0)
return np.ascontiguousarray(chw, dtype=np.float32)
def load_tensor_channel_names(tensor_path: Path) -> List[str]:
group_dir = tensor_path.parent.parent
meta_path = group_dir / "metas" / f"{tensor_path.stem}.json"
if not meta_path.is_file():
return []
meta = load_json(meta_path)
return list(meta.get("channels") or meta.get("input_channels") or [])
def collect_tensor_samples(root: Path, max_samples: int = 20, start_idx: int = 0) -> List[Path]:
tensor_paths: List[Path] = []
direct = root / "tensors"
if direct.is_dir():
tensor_paths.extend(sorted(direct.glob("*.npy")))
group_root = root / "group"
if group_root.is_dir():
for gdir in sorted(group_root.iterdir()):
tdir = gdir / "tensors"
if tdir.is_dir():
tensor_paths.extend(sorted(tdir.glob("*.npy")))
if not tensor_paths:
tensor_paths.extend(sorted(root.glob("**/tensors/*.npy")))
if not tensor_paths:
raise RuntimeError(f"Nenhum tensor .npy encontrado em: {root}")
start_idx = max(0, int(start_idx))
selected = tensor_paths[start_idx:]
if max_samples > 0:
selected = selected[:int(max_samples)]
return selected
def find_norm_stats(config: dict, config_dir: Path, save_dir: Path, explicit: Optional[str]) -> Optional[Path]:
if explicit:
return resolve_path(explicit, config_dir)
W, H = config.get("resolucao", [1024, 640])
dataset_path = config_dir / "dataset"
candidates = [
save_dir / "norm_stats.json",
dataset_path / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json",
config_dir / "backup" / config.get("modelo", "segformer_b1") / config.get("model_name", "test") / config.get("stats_source_tag", "stacked_raw5") / "norm_stats.json",
]
for p in candidates:
if p.is_file():
return p
return candidates[0]
def experiment_tag(config: dict, channels: int) -> str:
explicit = str(config.get("stats_source_tag", "")).strip()
if explicit:
return explicit
return f"{config.get('fusion_mode', 'stacked')}_raw{channels}"
def load_export_metadata(onnx_path: Path) -> Tuple[dict, Optional[Path]]:
meta_path = onnx_path.with_suffix(".export_meta.json")
if meta_path.is_file():
return load_json(meta_path), meta_path
return {}, None
def nanmean_list(arr: np.ndarray) -> List[Optional[float]]:
if arr.size == 0:
return []
out = []
for col in range(arr.shape[1]):
v = arr[:, col]
v = v[~np.isnan(v)]
out.append(None if v.size == 0 else float(np.mean(v)))
return out
def nanmin_list(arr: np.ndarray) -> List[Optional[float]]:
if arr.size == 0:
return []
out = []
for col in range(arr.shape[1]):
v = arr[:, col]
v = v[~np.isnan(v)]
out.append(None if v.size == 0 else float(np.min(v)))
return out
def resolve_model_artifact_paths(
args,
config: dict,
config_dir: Path,
channels: int,
) -> Tuple[Path, Path, str]:
"""
Resolve checkpoint e ONNX.
Se --checkpoint ou --onnx forem informados, usa os caminhos informados.
Se ficarem vazios, monta a partir do config:
backup/{modelo}/{model_name}/{stats_source_tag}/{ckpt_name}.pt
backup/{modelo}/{model_name}/{stats_source_tag}/{ckpt_name}.onnx
ckpt_name vem de:
config["ckpt_test"] ou "best_score"
"""
model = config.get("modelo", "segformer_b1")
model_name = config.get("model_name", "target_teached")
ckpt_name = config.get("ckpt_test", "best_score")
base_dir = config_dir / "backup" / model / model_name / experiment_tag(config, channels)
ckpt_filename = str(ckpt_name)
if not ckpt_filename.lower().endswith(".pt"):
ckpt_filename += ".pt"
onnx_filename = f"{Path(str(ckpt_name)).stem}.onnx"
if args.checkpoint:
checkpoint_path = resolve_path(args.checkpoint, config_dir)
else:
checkpoint_path = base_dir / ckpt_filename
if args.onnx:
onnx_path = resolve_path(args.onnx, config_dir)
else:
onnx_path = base_dir / onnx_filename
if checkpoint_path is None or not checkpoint_path.is_file():
raise FileNotFoundError(
f"Checkpoint não encontrado: {checkpoint_path}\n"
f"Dica: informe --checkpoint ou ajuste config['ckpt_test']."
)
if onnx_path is None or not onnx_path.is_file():
raise FileNotFoundError(
f"ONNX não encontrado: {onnx_path}\n"
f"Dica: informe --onnx ou ajuste config['ckpt_test']."
)
return checkpoint_path.resolve(), onnx_path.resolve(), str(ckpt_name)
# ============================================================
# PyTorch wrapper
# ============================================================
class MultiHeadTorchWrapper(nn.Module):
def __init__(
self,
model: nn.Module,
output_heads: List[str],
postprocess: str = "none",
):
super().__init__()
self.model = model
self.output_heads = list(output_heads)
self.postprocess = str(postprocess).lower()
if self.postprocess not in ("none", "resize_logits", "argmax_lowres", "argmax_fullres"):
raise RuntimeError(f"postprocess inválido: {self.postprocess}")
def forward(self, pixel_values: torch.Tensor):
outputs: Dict[str, torch.Tensor] = self.model(pixel_values=pixel_values)
result = {}
input_hw = pixel_values.shape[-2:]
for head_name in self.output_heads:
logits = outputs[head_name]
if self.postprocess in ("resize_logits", "argmax_fullres"):
if logits.shape[-2:] != input_hw:
logits = F.interpolate(
logits,
size=input_hw,
mode="bilinear",
align_corners=False,
)
if self.postprocess in ("argmax_lowres", "argmax_fullres"):
result[head_name] = torch.argmax(logits, dim=1).to(torch.uint8)
else:
result[head_name] = logits
return result
# ============================================================
# ONNX
# ============================================================
def create_onnx_session(
onnx_path: Path,
provider: str,
trt_home: Optional[str] = None,
trt_fp16: bool = True,
):
try:
import onnxruntime as ort
except ImportError:
raise ImportError(
"onnxruntime não está instalado. Instale com:\n"
" pip install onnxruntime-gpu\n"
"ou, para CPU:\n"
" pip install onnxruntime"
)
provider = provider.lower()
# ========================================================
# Windows/DLL helper para TensorRT
# ========================================================
if provider == "tensorrt":
trt_home = trt_home or os.environ.get("TRT_HOME", r"C:\dev\TensorRT-10.10.0.31")
dll_dirs = [
os.path.join(trt_home, "lib"),
os.path.join(trt_home, "bin"),
]
cuda_home = os.environ.get("CUDA_PATH")
if cuda_home:
dll_dirs.append(os.path.join(cuda_home, "bin"))
# Fallback comum para CUDA 12.4
dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.4\bin")
for dll_dir in dll_dirs:
if os.path.isdir(dll_dir):
try:
os.add_dll_directory(dll_dir)
print(f"[DLL] add_dll_directory: {dll_dir}")
except Exception as e:
print(f"[DLL][WARN] falha em {dll_dir}: {e}")
available = ort.get_available_providers()
print(f"[ONNX] providers disponíveis: {available}")
sess_options = ort.SessionOptions()
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
if provider == "cuda":
providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
elif provider == "cpu":
providers = ["CPUExecutionProvider"]
elif provider == "tensorrt":
cache_dir = onnx_path.parent / "trt_cache"
cache_dir.mkdir(parents=True, exist_ok=True)
trt_options = {
"device_id": 0,
"trt_fp16_enable": bool(trt_fp16),
# Cache para não reconstruir engine/timing toda vez.
"trt_engine_cache_enable": True,
"trt_engine_cache_path": str(cache_dir),
"trt_timing_cache_enable": True,
"trt_timing_cache_path": str(cache_dir),
# 4GB de workspace. Sua RTX 3070 lidou bem no benchmark.
"trt_max_workspace_size": 4 * 1024 * 1024 * 1024,
}
providers = [
("TensorrtExecutionProvider", trt_options),
"CUDAExecutionProvider",
"CPUExecutionProvider",
]
else:
raise RuntimeError(f"Provider desconhecido: {provider}")
requested_names = [
p[0] if isinstance(p, tuple) else p
for p in providers
]
providers_ok = [
p for p in providers
if (p[0] if isinstance(p, tuple) else p) in available
]
if not providers_ok:
raise RuntimeError(
f"Nenhum provider solicitado está disponível. "
f"Solicitado={requested_names}, disponível={available}"
)
sess = ort.InferenceSession(
str(onnx_path),
sess_options=sess_options,
providers=providers_ok,
)
active = sess.get_providers()
print(f"[ONNX] usando providers: {active}")
# Trava anti-burrice silenciosa: se pedir TensorRT, não pode cair para CPU/CUDA sem avisar.
if provider == "tensorrt" and "TensorrtExecutionProvider" not in active:
raise RuntimeError(
"TensorRTExecutionProvider foi solicitado, mas não ficou ativo. "
f"Providers ativos: {active}. "
"Provável causa: DLLs TensorRT fora do PATH/add_dll_directory, "
"versão incompatível ou fallback interno."
)
if provider == "cuda" and "CUDAExecutionProvider" not in active:
raise RuntimeError(
"CUDAExecutionProvider foi solicitado, mas não ficou ativo. "
f"Providers ativos: {active}."
)
return sess
def run_onnx(session, input_name: str, x_nchw: np.ndarray) -> Dict[str, np.ndarray]:
outputs = session.run(None, {input_name: x_nchw.astype(np.float32)})
output_names = [o.name for o in session.get_outputs()]
if len(outputs) != len(output_names):
raise RuntimeError("Quantidade de outputs ONNX inesperada.")
return {
name: arr.astype(np.float32)
for name, arr in zip(output_names, outputs)
}
def normalize_onnx_output_names(onnx_outputs: Dict[str, np.ndarray], output_heads: List[str]) -> Dict[str, np.ndarray]:
"""
Converte:
semantic_logits -> semantic
vegetation_logits -> vegetation
etc.
"""
out = {}
for head in output_heads:
candidates = [
head,
f"{head}_logits",
f"{head}_mask",
f"output_{head}",
]
found = None
for c in candidates:
if c in onnx_outputs:
found = c
break
if found is None:
# fallback por ordem caso nomes estejam diferentes
keys = list(onnx_outputs.keys())
idx = output_heads.index(head)
if idx < len(keys):
found = keys[idx]
if found is None:
raise RuntimeError(f"Não encontrei saída ONNX para head={head}. Outputs={list(onnx_outputs.keys())}")
out[head] = onnx_outputs[found]
return out
# ============================================================
# Main
# ============================================================
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", default="config.json")
parser.add_argument("--checkpoint", default="")
parser.add_argument("--onnx", default="")
parser.add_argument("--train-script", default="_8_train_multihead.py")
parser.add_argument("--labelmap", default="dataset/labelmap.txt")
parser.add_argument("--split_folder", default="val", choices=["train", "val", "test"])
parser.add_argument("--root_override", default=None)
parser.add_argument("--norm_stats", default=None)
parser.add_argument("--max_samples", type=int, default=20)
parser.add_argument("--start_idx", type=int, default=0)
parser.add_argument("--device", default="cuda", choices=["cuda", "cpu"])
parser.add_argument("--onnx_provider", default="cuda", choices=["cuda", "cpu", "tensorrt"])
parser.add_argument(
"--trt_home",
default="C:\\dev\\TensorRT-10.10.0.31",
help="Pasta raiz do TensorRT. Ex: C:\\dev\\TensorRT-10.10.0.31. Se omitido, usa TRT_HOME ou fallback padrão.",
)
parser.add_argument(
"--trt_no_fp16",
action="store_true",
help="Desativa FP16 no TensorRT. Normalmente NÃO usar; deixamos FP16 ligado.",
)
parser.add_argument(
"--torch_no_amp",
action="store_true",
help="Desativa AMP no PyTorch. Recomendado para comparação mais rígida contra ONNX FP32.",
)
parser.add_argument(
"--resize_torch_to_input",
action="store_true",
help="Força saída PyTorch redimensionada para HxW antes de comparar.",
)
parser.add_argument(
"--onnx_has_norm",
action="store_true",
help="Força ONNX com normalização interna; normalmente detectado pelo export_meta.",
)
parser.add_argument(
"--onnx_no_norm",
action="store_true",
help="Força ONNX sem normalização interna; normalmente detectado pelo export_meta.",
)
parser.add_argument(
"--onnx_output_kind",
default="auto",
choices=["auto", "logits", "mask"],
help="Tipo de saída; auto usa o metadata gerado pelo exportador.",
)
parser.add_argument(
"--compare_at_input_size",
action="store_true",
help="Redimensiona ambos os outputs para HxW antes de comparar.",
)
parser.add_argument(
"--save_report",
default=None,
help="Caminho do JSON de relatório. Se omitido, salva ao lado do ONNX.",
)
args = parser.parse_args()
config_path = resolve_path(args.config, Path.cwd())
if config_path is None or not config_path.is_file():
raise FileNotFoundError(f"Config não encontrado: {config_path}")
config_dir = config_path.parent
config = load_json(config_path)
train_script_path = resolve_path(args.train_script, config_dir)
labelmap_path = resolve_path(args.labelmap, config_dir)
if train_script_path is None or not train_script_path.is_file():
raise FileNotFoundError(f"Train script não encontrado: {train_script_path}")
if labelmap_path is None or not labelmap_path.is_file():
raise FileNotFoundError(f"Labelmap não encontrado: {labelmap_path}")
train_mod = import_train_module(train_script_path)
W, H = config.get("resolucao", [1024, 640])
W = int(W)
H = int(H)
backbone = config.get("backbone", "nvidia/mit-b1")
input_channel_names = train_mod.get_input_channel_names(config)
input_channel_indices = train_mod.get_input_channel_indices(config)
channels = len(input_channel_names)
checkpoint_path, onnx_path, ckpt_name = resolve_model_artifact_paths(
args=args,
config=config,
config_dir=config_dir,
channels=channels,
)
export_meta, export_meta_path = load_export_metadata(onnx_path)
exported_names = [
str(x).strip().upper()
for x in export_meta.get("input_channel_names", [])
]
if exported_names and exported_names != input_channel_names:
raise RuntimeError(
f"ONNX incompatível: canais exportados={exported_names}, "
f"config atual={input_channel_names}."
)
exported_heads = list(export_meta.get("heads", []) or [])
exported_postprocess = str(export_meta.get("postprocess", "none")).lower()
if exported_postprocess == "none" and bool(export_meta.get("resize_to_input", False)):
exported_postprocess = "resize_logits"
if not export_meta:
if args.onnx_output_kind == "mask":
exported_postprocess = (
"argmax_fullres" if args.resize_torch_to_input else "argmax_lowres"
)
elif args.resize_torch_to_input:
exported_postprocess = "resize_logits"
if exported_postprocess not in ("none", "resize_logits", "argmax_lowres", "argmax_fullres"):
raise RuntimeError(f"postprocess inválido no export_meta: {exported_postprocess}")
if args.onnx_has_norm and args.onnx_no_norm:
raise RuntimeError("Use apenas uma opção: --onnx_has_norm ou --onnx_no_norm.")
input_contract = export_meta.get("input_contract", {}) or {}
detected_onnx_has_norm = bool(
input_contract.get("normalization_embedded", export_meta.get("include_norm", False))
)
onnx_has_norm = (
True if args.onnx_has_norm
else False if args.onnx_no_norm
else detected_onnx_has_norm
)
detected_output_kind = (
"mask" if exported_postprocess in ("argmax_lowres", "argmax_fullres") else "logits"
)
onnx_output_kind = (
detected_output_kind if args.onnx_output_kind == "auto" else args.onnx_output_kind
)
if export_meta and onnx_output_kind != detected_output_kind:
raise RuntimeError(
f"--onnx_output_kind={onnx_output_kind} contradiz export_meta: "
f"postprocess={exported_postprocess} ({detected_output_kind})."
)
semantic_id2label, semantic_label2id, ignore_from_labelmap = train_mod.load_labelmap(
str(labelmap_path)
)
heads_config = train_mod.build_heads_config(
config,
ignore_index=int(ignore_from_labelmap)
)
heads_config["semantic"]["num_classes"] = int(len(semantic_id2label))
heads_config["semantic"]["ignore_index"] = int(ignore_from_labelmap)
output_heads = list(heads_config.keys())
if exported_heads and exported_heads != output_heads:
raise RuntimeError(
f"Heads do ONNX={exported_heads} diferem do config/modelo={output_heads}."
)
save_dir = (
config_dir
/ "backup"
/ config.get("modelo", "segformer_b1")
/ config.get("model_name", "test")
/ experiment_tag(config, channels)
)
norm_stats_path = find_norm_stats(
config=config,
config_dir=config_dir,
save_dir=save_dir,
explicit=args.norm_stats,
)
mean, std, norm_stats_used = load_norm_stats(
norm_stats_path,
channel_names=input_channel_names,
)
if args.root_override:
root = resolve_path(args.root_override, Path.cwd())
else:
root = (config_dir / "dataset" / "split" / args.split_folder).resolve()
if root is None or not root.is_dir():
raise FileNotFoundError(f"Root de dados não encontrado: {root}")
samples = collect_tensor_samples(
root=root,
max_samples=args.max_samples,
start_idx=args.start_idx,
)
use_cuda = args.device == "cuda" and torch.cuda.is_available()
device = torch.device("cuda" if use_cuda else "cpu")
if args.device == "cuda" and not torch.cuda.is_available():
print("[WARN] CUDA indisponível. Usando CPU no PyTorch.")
print("==========================================")
print("Validate PyTorch vs ONNX")
print(f"Config : {config_path}")
print(f"Checkpoint : {checkpoint_path}")
print(f"ONNX : {onnx_path}")
print(f"Export meta : {export_meta_path or 'ausente (modo legado)'}")
print(f"Root : {root}")
print(f"Samples : {len(samples)}")
print(f"Backbone : {backbone}")
print(f"Input shape : [1, {channels}, {H}, {W}]")
print(f"Channels : {input_channel_names} idx={input_channel_indices}")
print(f"Heads : {output_heads}")
print(f"Device : {device}")
print(f"ONNX provider: {args.onnx_provider}")
print(f"Torch AMP : {not args.torch_no_amp and device.type == 'cuda'}")
print(f"ONNX norm : {'embutida' if onnx_has_norm else 'externa'}")
print(f"Postprocess : {exported_postprocess}")
print(f"Compare HxW : {args.compare_at_input_size}")
print("==========================================")
print("[MODEL] Montando PyTorch...")
model = train_mod.build_model(
backbone=backbone,
input_channel_names=input_channel_names,
heads_config=heads_config,
semantic_id2label=semantic_id2label,
semantic_label2id=semantic_label2id,
)
ckpt = train_mod.load_checkpoint(
str(checkpoint_path),
model,
map_location="cpu",
)
saved_contract = ckpt.get("model_contract", {}) or {}
saved_names = [str(x).strip().upper() for x in saved_contract.get("input_channel_names", [])]
legacy_names = SOURCE_CHANNEL_ORDER[:channels]
if not saved_names and input_channel_names != legacy_names:
raise RuntimeError(
"Checkpoint sem contrato nominal não pode ser validado para "
f"input_channels={input_channel_names}."
)
model.to(device)
model.eval()
torch_wrapper = MultiHeadTorchWrapper(
model=model,
output_heads=output_heads,
postprocess=exported_postprocess,
).to(device)
torch_wrapper.eval()
print("[ONNX] Carregando sessão...")
onnx_session = create_onnx_session(
onnx_path=onnx_path,
provider=args.onnx_provider,
trt_home=args.trt_home,
trt_fp16=not args.trt_no_fp16,
)
onnx_input = onnx_session.get_inputs()[0]
onnx_input_name = onnx_input.name
onnx_input_shape = onnx_input.shape
if len(onnx_input_shape) != 4:
raise RuntimeError(f"Entrada ONNX inesperada: {onnx_input_name} shape={onnx_input_shape}")
onnx_channels = onnx_input_shape[1]
if isinstance(onnx_channels, int) and onnx_channels != channels:
raise RuntimeError(
f"ONNX espera {onnx_channels} canais, config atual possui {channels}: "
f"{input_channel_names}"
)
print(f"[ONNX] input: {onnx_input_name} shape={onnx_input_shape}")
print(f"[ONNX] outputs: {[o.name for o in onnx_session.get_outputs()]}")
per_head_accum = {
h: {
"n": 0,
"logits_abs_mean": [],
"logits_abs_max": [],
"prob_abs_mean": [],
"prob_abs_max": [],
"argmax_equal_ratio": [],
"pred_miou_torch_vs_onnx": [],
"pred_iou_per_class": [],
}
for h in output_heads
}
sample_reports = []
for i, tensor_path in enumerate(samples):
saved_channel_names = load_tensor_channel_names(tensor_path)
chw01 = load_tensor(
tensor_path,
requested_channels=input_channel_names,
saved_channel_names=saved_channel_names,
)
# Garante resolução do contrato.
if chw01.shape[-2:] != (H, W):
hwc = np.transpose(chw01, (1, 2, 0))
hwc = cv2.resize(hwc, (W, H), interpolation=cv2.INTER_LINEAR)
chw01 = np.transpose(hwc, (2, 0, 1)).astype(np.float32)
chw_norm = normalize_numpy_chw(chw01, mean=mean, std=std)
# PyTorch continua recebendo normalizado, porque o modelo PyTorch puro espera isso.
x_torch_np = np.expand_dims(chw_norm, axis=0).astype(np.float32)
# ONNX novo com include_norm recebe 0..1 cru.
if onnx_has_norm:
x_onnx_np = np.expand_dims(chw01, axis=0).astype(np.float32)
else:
x_onnx_np = x_torch_np
x_torch = torch.from_numpy(x_torch_np).to(device, non_blocking=True)
if device.type == "cuda":
torch.cuda.synchronize()
t0 = time.perf_counter()
with torch.inference_mode():
with torch.autocast(
device_type="cuda",
dtype=torch.float16,
enabled=(not args.torch_no_amp and device.type == "cuda"),
):
torch_outputs_t = torch_wrapper(x_torch)
if device.type == "cuda":
torch.cuda.synchronize()
torch_ms = (time.perf_counter() - t0) * 1000.0
torch_outputs = {
h: v.detach().float().cpu().numpy()
for h, v in torch_outputs_t.items()
}
t0 = time.perf_counter()
onnx_raw_outputs = run_onnx(onnx_session, onnx_input_name, x_onnx_np)
onnx_ms = (time.perf_counter() - t0) * 1000.0
onnx_outputs = normalize_onnx_output_names(
onnx_raw_outputs,
output_heads=output_heads,
)
report_item = {
"idx": i,
"tensor": str(tensor_path),
"torch_ms": float(torch_ms),
"onnx_ms": float(onnx_ms),
"heads": {},
}
print(f"\n[{i + 1:03d}/{len(samples):03d}] {tensor_path.name} | torch={torch_ms:.2f}ms | onnx={onnx_ms:.2f}ms")
for head in output_heads:
pt = torch_outputs[head]
ox = onnx_outputs[head]
if onnx_output_kind == "mask":
# Esperado:
# PT : [1,H,W]
# ONNX : [1,H,W]
pt_mask = np.asarray(pt).astype(np.uint8)
ox_mask = np.asarray(ox).astype(np.uint8)
if pt_mask.ndim == 3:
pt_mask = pt_mask[0]
if ox_mask.ndim == 3:
ox_mask = ox_mask[0]
if pt_mask.shape != ox_mask.shape:
ox_mask = cv2.resize(
ox_mask,
(pt_mask.shape[1], pt_mask.shape[0]),
interpolation=cv2.INTER_NEAREST,
)
equal_ratio = float(np.mean(pt_mask == ox_mask))
num_classes = int(heads_config[head]["num_classes"])
iou_per_class, miou, present_classes = compute_mask_iou_between_preds(
pt_mask,
ox_mask,
num_classes=num_classes,
)
head_report = {
"torch_shape": list(np.asarray(pt).shape),
"onnx_shape": list(np.asarray(ox).shape),
"compare_shape": list(pt_mask.shape),
"logits_abs_mean": None,
"logits_abs_max": None,
"prob_abs_mean": None,
"prob_abs_max": None,
"argmax_equal_ratio": float(equal_ratio),
"pred_miou_torch_vs_onnx": float(miou),
"pred_iou_per_class": [
None if x is None else float(x)
for x in iou_per_class
],
"pred_present_classes": [int(x) for x in present_classes],
}
print(
f" {head:<10} "
f"shape PT={tuple(np.asarray(pt).shape)} ONNX={tuple(np.asarray(ox).shape)} "
f"CMP={tuple(pt_mask.shape)} | "
f"mask_equal={equal_ratio * 100:.3f}% "
f"mIoU={miou:.6f}"
)
else:
pt = np.asarray(pt).astype(np.float32)
ox = np.asarray(ox).astype(np.float32)
compare_hw = (H, W) if args.compare_at_input_size else None
if pt.shape != ox.shape:
compare_hw = (H, W)
if compare_hw is not None:
pt_cmp = resize_logits_np_nchw(pt, compare_hw)
ox_cmp = resize_logits_np_nchw(ox, compare_hw)
else:
pt_cmp = pt
ox_cmp = ox
if pt_cmp.shape != ox_cmp.shape:
raise RuntimeError(
f"Shape incompatível na head {head}: "
f"torch={pt_cmp.shape}, onnx={ox_cmp.shape}"
)
diff_logits = np.abs(pt_cmp - ox_cmp)
prob_pt = softmax_np(pt_cmp, axis=1)
prob_ox = softmax_np(ox_cmp, axis=1)
diff_prob = np.abs(prob_pt - prob_ox)
pred_pt = np.argmax(prob_pt, axis=1)[0].astype(np.uint8)
pred_ox = np.argmax(prob_ox, axis=1)[0].astype(np.uint8)
equal_ratio = float(np.mean(pred_pt == pred_ox))
num_classes = int(heads_config[head]["num_classes"])
iou_per_class, miou, present_classes = compute_mask_iou_between_preds(
pred_pt,
pred_ox,
num_classes=num_classes,
)
head_report = {
"torch_shape": list(pt.shape),
"onnx_shape": list(ox.shape),
"compare_shape": list(pt_cmp.shape),
"logits_abs_mean": float(diff_logits.mean()),
"logits_abs_max": float(diff_logits.max()),
"prob_abs_mean": float(diff_prob.mean()),
"prob_abs_max": float(diff_prob.max()),
"argmax_equal_ratio": float(equal_ratio),
"pred_miou_torch_vs_onnx": float(miou),
"pred_iou_per_class": [
None if x is None else float(x)
for x in iou_per_class
],
"pred_present_classes": [int(x) for x in present_classes],
}
print(
f" {head:<10} "
f"shape PT={tuple(pt.shape)} ONNX={tuple(ox.shape)} CMP={tuple(pt_cmp.shape)} | "
f"logit_mean={head_report['logits_abs_mean']:.6g} "
f"prob_mean={head_report['prob_abs_mean']:.6g} "
f"argmax_equal={equal_ratio * 100:.3f}% "
f"mIoU={miou:.6f}"
)
report_item["heads"][head] = head_report
acc = per_head_accum[head]
acc["n"] += 1
if head_report["logits_abs_mean"] is not None:
acc["logits_abs_mean"].append(head_report["logits_abs_mean"])
if head_report["logits_abs_max"] is not None:
acc["logits_abs_max"].append(head_report["logits_abs_max"])
if head_report["prob_abs_mean"] is not None:
acc["prob_abs_mean"].append(head_report["prob_abs_mean"])
if head_report["prob_abs_max"] is not None:
acc["prob_abs_max"].append(head_report["prob_abs_max"])
acc["argmax_equal_ratio"].append(head_report["argmax_equal_ratio"])
acc["pred_miou_torch_vs_onnx"].append(head_report["pred_miou_torch_vs_onnx"])
acc["pred_iou_per_class"].append(head_report["pred_iou_per_class"])
#print(
# f" {head:<10} "
# f"shape PT={tuple(pt.shape)} ONNX={tuple(ox.shape)} CMP={tuple(pt_cmp.shape)} | "
# f"logit_mean={head_report['logits_abs_mean']:.6g} "
# f"prob_mean={head_report['prob_abs_mean']:.6g} "
# f"argmax_equal={equal_ratio * 100:.3f}% "
# f"mIoU={miou:.6f}"
#)
sample_reports.append(report_item)
summary = {
"config": str(config_path),
"checkpoint": str(checkpoint_path),
"ckpt_name": ckpt_name,
"onnx": str(onnx_path),
"export_meta": str(export_meta_path) if export_meta_path is not None else None,
"root": str(root),
"samples": len(samples),
"input_shape": [1, channels, H, W],
"input_channel_names": input_channel_names,
"input_channel_indices": input_channel_indices,
"heads": output_heads,
"norm_stats_used": norm_stats_used,
"torch_amp": bool(not args.torch_no_amp and device.type == "cuda"),
"onnx_provider": args.onnx_provider,
"onnx_has_norm": bool(onnx_has_norm),
"onnx_output_kind": onnx_output_kind,
"onnx_postprocess": exported_postprocess,
"trt_home": args.trt_home or os.environ.get("TRT_HOME", None),
"trt_fp16": bool(not args.trt_no_fp16),
"compare_at_input_size": bool(args.compare_at_input_size),
"per_head": {},
"sample_reports": sample_reports,
}
print("\n========== RESUMO ==========")
for head, acc in per_head_accum.items():
if acc["n"] <= 0:
continue
ious_raw = acc["pred_iou_per_class"]
if ious_raw:
ious_arr = np.array(
[
[np.nan if x is None else float(x) for x in row]
for row in ious_raw
],
dtype=np.float64,
)
else:
ious_arr = np.empty((0, 0), dtype=np.float64)
def safe_mean(values):
return None if not values else float(np.mean(values))
def safe_max(values):
return None if not values else float(np.max(values))
head_summary = {
"n": int(acc["n"]),
"logits_abs_mean_avg": safe_mean(acc["logits_abs_mean"]),
"logits_abs_mean_max": safe_max(acc["logits_abs_mean"]),
"logits_abs_max_avg": safe_mean(acc["logits_abs_max"]),
"logits_abs_max_max": safe_max(acc["logits_abs_max"]),
"prob_abs_mean_avg": safe_mean(acc["prob_abs_mean"]),
"prob_abs_mean_max": safe_max(acc["prob_abs_mean"]),
"prob_abs_max_avg": safe_mean(acc["prob_abs_max"]),
"prob_abs_max_max": safe_max(acc["prob_abs_max"]),
"argmax_equal_ratio_avg": float(np.mean(acc["argmax_equal_ratio"])),
"argmax_equal_ratio_min": float(np.min(acc["argmax_equal_ratio"])),
"pred_miou_avg": float(np.mean(acc["pred_miou_torch_vs_onnx"])),
"pred_miou_min": float(np.min(acc["pred_miou_torch_vs_onnx"])),
"pred_iou_per_class_avg": nanmean_list(ious_arr),
"pred_iou_per_class_min": nanmin_list(ious_arr),
}
summary["per_head"][head] = head_summary
def fmt_iou_list(values):
return [
None if x is None else round(float(x), 6)
for x in values
]
def fmt_optional(v, casas=8):
if v is None:
return "N/A"
return f"{float(v):.{casas}f}"
print(f"\n[{head}]")
print(f" logits_abs_mean avg : {fmt_optional(head_summary['logits_abs_mean_avg'])}")
print(f" prob_abs_mean avg : {fmt_optional(head_summary['prob_abs_mean_avg'])}")
print(f" argmax_equal avg : {head_summary['argmax_equal_ratio_avg'] * 100:.4f}%")
print(f" argmax_equal min : {head_summary['argmax_equal_ratio_min'] * 100:.4f}%")
print(f" pred_mIoU avg : {head_summary['pred_miou_avg']:.8f}")
print(f" pred_mIoU min : {head_summary['pred_miou_min']:.8f}")
print(f" IoU/classes avg : {fmt_iou_list(head_summary['pred_iou_per_class_avg'])}")
print(f" IoU/classes min : {fmt_iou_list(head_summary['pred_iou_per_class_min'])}")
if args.save_report:
report_path = resolve_path(args.save_report, config_dir)
else:
suffix = f".validate_{args.onnx_provider}_report.json"
report_path = onnx_path.with_suffix(suffix)
save_json(report_path, summary)
print(f"\n[OK] Relatório salvo em: {report_path}")
print("\nValidação finalizada.")
if __name__ == "__main__":
main()