diff --git a/Python/OAK/datasets/_10_export_onnx.py b/Python/OAK/datasets/_10_export_onnx.py new file mode 100644 index 000000000..accf28440 --- /dev/null +++ b/Python/OAK/datasets/_10_export_onnx.py @@ -0,0 +1,841 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +_10_export_visual_onnx.py + +Exporta o checkpoint PyTorch do Visual Worker para ONNX. + +Contrato esperado do modelo visual: + - Backbone SegFormer Hugging Face, ex: nvidia/mit-b0 + - Entrada RGB: [N, 3, H, W], float32 + - Cabeça 1: segmentação semântica do corredor + - Cabeça 2: classificação global/status do corredor + +Suporta: + - ONNX com logits crus de segmentação + - ONNX com logits de segmentação redimensionados para HxW + - ONNX com máscara argmax de segmentação + - ONNX com label_logits ou label_probs + - Normalização embutida no grafo: x = (x - mean) / std + +Exemplos: + +# ONNX recomendado para runtime: recebe RGB 0..1 e já normaliza internamente. +python _10_export_visual_onnx.py ^ + --config config.json ^ + --checkpoint oak-d/backup/segformer_b0/nav_mit_dual_label/best_label.pt ^ + --out oak-d/backup/segformer_b0/nav_mit_dual_label/best_label.onnx ^ + --device cuda ^ + --include-norm ^ + --semantic-postprocess resize_logits ^ + --label-postprocess probs + +# ONNX cru: runtime precisa enviar tensor já normalizado. +python _10_export_visual_onnx.py ^ + --config config.json ^ + --device cuda ^ + --semantic-postprocess none ^ + --label-postprocess logits +""" + +from __future__ import annotations + +import json +import argparse +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import SegformerForSemanticSegmentation + + +# ============================================================ +# 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 ensure_dir(path: str | Path): + Path(path).mkdir(parents=True, exist_ok=True) + + +def resolve_path(path_like: Optional[str], base: Optional[Path] = None) -> Optional[Path]: + if path_like is None or str(path_like).strip() == "": + return None + + p = Path(path_like) + if p.is_absolute(): + return p + + if base is None: + base = Path.cwd() + + return (base / p).resolve() + + +def _try_int(text: str) -> Optional[int]: + try: + return int(str(text).strip()) + except Exception: + return None + + +def clean_label_name(raw_name: str) -> str: + """ + Limpa nomes vindos de labelmap no estilo: + naonavegavel:128,0,0:: + navegavel:0,128,0:: + + Mantém apenas o nome lógico da classe. + """ + name = str(raw_name).strip() + + # Remove sufixos visuais comuns do labelmap: classe:R,G,B:: + if "::" in name: + name = name.split("::", 1)[0].strip() + + if ":" in name: + left, right = name.split(":", 1) + right_clean = right.replace(",", "").replace(" ", "") + if right_clean.isdigit(): + name = left.strip() + + return name.strip() + + +def load_labelmap(labelmap_path: Path) -> Tuple[Dict[int, str], Dict[str, int], int]: + """ + Parser tolerante para labelmap.txt. + + Aceita formatos comuns: + navegavel + 0:navegavel + 0 navegavel + 0,navegavel + navegavel:0 + navegavel:0,128,0:: + + Retorna: + id2label, label2id, ignore_index + """ + if not labelmap_path.is_file(): + raise FileNotFoundError(f"Labelmap não encontrado: {labelmap_path}") + + id2label: Dict[int, str] = {} + ignore_index = 255 + next_id = 0 + + with labelmap_path.open("r", encoding="utf-8") as f: + for raw_line in f: + line = raw_line.strip() + if not line or line.startswith("#"): + continue + + lower = line.lower() + if lower.startswith("ignore") or lower.startswith("ignore_index"): + for sep in ("=", ":", ",", " "): + if sep in line: + maybe = _try_int(line.split(sep)[-1]) + if maybe is not None: + ignore_index = maybe + break + continue + + cls_id: Optional[int] = None + cls_name: Optional[str] = None + + # Caso CVAT/labelmap visual: nome:R,G,B:: + # Exemplo: navegavel:0,128,0:: + if "::" in line and ":" in line: + before = line.split("::", 1)[0].strip() + maybe_name = before.split(":", 1)[0].strip() + if maybe_name: + cls_id = next_id + cls_name = maybe_name + + if cls_id is None: + for sep in (":", ",", " ", " "): + if sep in line: + parts = [p.strip() for p in line.split(sep) if p.strip()] + if len(parts) >= 2: + left_id = _try_int(parts[0]) + right_id = _try_int(parts[-1]) + + if left_id is not None: + cls_id = left_id + cls_name = sep.join(parts[1:]).strip() if sep in (":", ",") else " ".join(parts[1:]).strip() + break + + if right_id is not None: + cls_id = right_id + cls_name = sep.join(parts[:-1]).strip() if sep in (":", ",") else " ".join(parts[:-1]).strip() + break + + if cls_id is None: + cls_id = next_id + cls_name = line + + if cls_name is None or cls_name == "": + raise RuntimeError(f"Linha inválida no labelmap: {raw_line!r}") + + id2label[int(cls_id)] = clean_label_name(str(cls_name)) + next_id = max(next_id, int(cls_id) + 1) + + if not id2label: + raise RuntimeError(f"Labelmap vazio ou inválido: {labelmap_path}") + + # Garante ids contínuos para o SegFormer. + ids_sorted = sorted(id2label.keys()) + if ids_sorted != list(range(len(ids_sorted))): + remap = {old_id: new_id for new_id, old_id in enumerate(ids_sorted)} + id2label = {remap[old_id]: name for old_id, name in id2label.items()} + + label2id = {name: idx for idx, name in id2label.items()} + return id2label, label2id, ignore_index + + +def get_visual_mode(config: dict) -> str: + use_mask2 = bool(config.get("dual_head_mask", config.get("dual_head", False))) + use_label = bool(config.get("dual_head_label", False)) + + if use_mask2 and use_label: + raise RuntimeError("Config inválido: dual_head_mask e dual_head_label ativos juntos.") + + if use_label: + return "label" + if use_mask2: + return "mask2" + return "single" + + +def get_save_suffix(mode: str) -> str: + if mode == "single": + return "_single" + if mode == "mask2": + return "_dual_mask" + if mode == "label": + return "_dual_label" + raise RuntimeError(f"Modo desconhecido: {mode}") + + +def resolve_default_paths(args, config: dict, config_dir: Path) -> Tuple[Path, Path, str, str, Path]: + """ + Resolve checkpoint e ONNX padrão seguindo o padrão do treino visual: + {camera}/backup/{modelo}/{model_name}_{suffix}/{ckpt_name}.pt + """ + camera = str(config.get("camera", "oak-d")) + model_key = str(config.get("modelo", "segformer_b0")) + model_name = str(config.get("model_name", "visual")) + mode = get_visual_mode(config) + suffix = get_save_suffix(mode) + + save_dir = config_dir / camera / "backup" / model_key / f"{model_name}{suffix}" + + if mode == "label": + default_ckpt_name = "best_label" + elif mode == "mask2": + default_ckpt_name = "best_mask2" + else: + default_ckpt_name = "best_main" + + ckpt_name = str(config.get("ckpt_test", default_ckpt_name)) + + checkpoint_path = resolve_path(args.checkpoint, Path.cwd()) + if checkpoint_path is None: + checkpoint_path = save_dir / f"{ckpt_name}.pt" + + out_path = resolve_path(args.out, Path.cwd()) + if out_path is None: + out_path = save_dir / f"{ckpt_name}.onnx" + + return checkpoint_path.resolve(), out_path.resolve(), ckpt_name, mode, save_dir.resolve() + + +def resolve_labelmap_path(args, config: dict, config_dir: Path) -> Path: + explicit = resolve_path(args.labelmap, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + camera = str(config.get("camera", "oak-d")) + return (config_dir / camera / "dataset" / "labelmap.txt").resolve() + + +def resolve_norm_stats_path(args, config: dict, config_dir: Path, save_dir: Path) -> Optional[Path]: + explicit = resolve_path(args.norm_stats, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + # Prioridade 1: norm_stats salvo junto ao treino/checkpoint. + p = save_dir / "norm_stats.json" + if p.is_file(): + return p.resolve() + + # Prioridade 2: dataset normalizado na resolução do contrato. + W, H = config.get("resolucao", [1024, 640]) + camera = str(config.get("camera", "oak-d")) + p = config_dir / camera / "dataset" / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json" + if p.is_file(): + return p.resolve() + + # Prioridade 3: deixa explícito no erro quando --include-norm for usado. + return p.resolve() + + +def load_rgb_norm_stats(path: Path) -> Tuple[List[float], List[float], List[str]]: + if not path.is_file(): + raise FileNotFoundError(f"norm_stats não encontrado: {path}") + + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + if names: + # Visual worker espera RGB. Aceita stats com canais extras, desde que R,G,B existam. + name_to_idx = {str(n).upper(): i for i, n in enumerate(names)} + required = ["R", "G", "B"] + missing = [ch for ch in required if ch not in name_to_idx] + if missing: + raise RuntimeError( + f"norm_stats incompatível: faltam canais {missing}. " + f"channels={names} path={path}" + ) + idx = [name_to_idx[ch] for ch in required] + mean_sel = [float(mean[i]) for i in idx] + std_sel = [float(std[i]) for i in idx] + names_sel = required + else: + if len(mean) < 3 or len(std) < 3: + raise RuntimeError(f"norm_stats precisa de pelo menos 3 valores RGB: {path}") + mean_sel = [float(mean[i]) for i in range(3)] + std_sel = [float(std[i]) for i in range(3)] + names_sel = ["R", "G", "B"] + + return mean_sel, std_sel, names_sel + + +def resolve_label_classes(config: dict, ckpt: Optional[dict] = None) -> Tuple[Dict[int, str], int]: + """Resolve classes da cabeça de status/corredor.""" + label_classes = config.get("label_classes", None) + + if label_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(label_classes)} + return label_name_by_id, len(label_name_by_id) + + if ckpt is not None: + extra = ckpt.get("extra", {}) if isinstance(ckpt, dict) else {} + maybe = extra.get("label_name_by_id", None) + if isinstance(maybe, dict) and maybe: + label_name_by_id = {int(k): str(v) for k, v in maybe.items()} + return label_name_by_id, max(label_name_by_id.keys()) + 1 + + maybe_config = extra.get("config", {}) if isinstance(extra, dict) else {} + maybe_classes = maybe_config.get("label_classes", None) if isinstance(maybe_config, dict) else None + if maybe_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(maybe_classes)} + return label_name_by_id, len(label_name_by_id) + + raise RuntimeError( + "Não consegui resolver label_classes. " + "Adicione config['label_classes'] com a lista de status do corredor, " + "ou use um checkpoint que tenha extra['label_name_by_id']." + ) + + +# ============================================================ +# Cabeça auxiliar igual ao treino visual +# ============================================================ + + +class LabelHead(nn.Module): + """Head de classificação global do frame/status do corredor.""" + + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + hidden: int = 256, + dropout: float = 0.2, + ): + super().__init__() + in_ch = int(feat_ch) + int(num_seg_classes) + self.in_ch = in_ch + self.pool = nn.AdaptiveAvgPool2d((1, 1)) + self.net = nn.Sequential( + nn.Linear(in_ch, hidden), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + nn.Linear(hidden, num_label_classes), + ) + + def forward(self, feat: torch.Tensor, logits_seg: torch.Tensor) -> torch.Tensor: + # Interpola sempre para evitar branch Python durante o trace ONNX. + # Se o tamanho já for igual, o resultado é equivalente e o grafo fica estável. + feat = F.interpolate( + feat, + size=logits_seg.shape[-2:], + mode="bilinear", + align_corners=False, + ) + x = torch.cat([feat, logits_seg], dim=1) + x = self.pool(x).flatten(1) + return self.net(x) + + +def get_last_feat(out, logits: torch.Tensor) -> torch.Tensor: + if hasattr(out, "hidden_states") and out.hidden_states is not None: + feat = out.hidden_states[-1] + else: + feat = logits + + # Interpola sempre para evitar TracerWarning por comparação de shapes. + # O contrato do export é resolução fixa, então isso não muda o comportamento útil. + feat = F.interpolate( + feat, + size=logits.shape[-2:], + mode="bilinear", + align_corners=False, + ) + return feat + + +# ============================================================ +# Modelo composto Visual +# ============================================================ + + +class VisualSegformerDualLabel(nn.Module): + """ + Modelo composto para export: + - base_model: SegformerForSemanticSegmentation + - label_head: LabelHead treinada junto + + forward retorna dict: + semantic: [N, Cseg, h, w] + label : [N, Cstatus] + """ + + def __init__(self, base_model: nn.Module, label_head: nn.Module): + super().__init__() + self.base_model = base_model + self.label_head = label_head + + def forward(self, pixel_values: torch.Tensor) -> Dict[str, torch.Tensor]: + out = self.base_model(pixel_values=pixel_values) + logits_seg = out.logits + feat = get_last_feat(out, logits_seg) + logits_label = self.label_head(feat, logits_seg) + return { + "semantic": logits_seg, + "label": logits_label, + } + + +class VisualOnnxWrapper(nn.Module): + """ + Wrapper final exportável para ONNX. + + semantic_postprocess: + - none -> semantic_logits em baixa resolução do decode_head + - resize_logits -> semantic_logits em HxW da entrada + - argmax_lowres -> semantic_mask em baixa resolução + - argmax_fullres-> semantic_mask em HxW da entrada + + label_postprocess: + - logits -> label_logits + - probs -> label_probs via softmax + """ + + def __init__( + self, + model: nn.Module, + include_norm: bool = False, + norm_mean: Optional[List[float]] = None, + norm_std: Optional[List[float]] = None, + semantic_postprocess: str = "none", + label_postprocess: str = "logits", + ): + super().__init__() + self.model = model + self.include_norm = bool(include_norm) + self.semantic_postprocess = str(semantic_postprocess).lower() + self.label_postprocess = str(label_postprocess).lower() + + if self.semantic_postprocess not in ("none", "resize_logits", "argmax_lowres", "argmax_fullres"): + raise RuntimeError(f"semantic_postprocess inválido: {self.semantic_postprocess}") + + if self.label_postprocess not in ("logits", "probs"): + raise RuntimeError(f"label_postprocess inválido: {self.label_postprocess}") + + if self.include_norm: + if norm_mean is None or norm_std is None: + raise RuntimeError("include_norm=True requer norm_mean/norm_std") + mean_t = torch.tensor(norm_mean, dtype=torch.float32).view(1, 3, 1, 1) + std_t = torch.tensor(norm_std, dtype=torch.float32).view(1, 3, 1, 1) + self.register_buffer("norm_mean", mean_t) + self.register_buffer("norm_std", std_t) + else: + self.register_buffer("norm_mean", torch.empty(0)) + self.register_buffer("norm_std", torch.empty(0)) + + def forward(self, pixel_values: torch.Tensor): + input_hw = pixel_values.shape[-2:] + + x = pixel_values + if self.include_norm: + x = (x - self.norm_mean) / torch.clamp(self.norm_std, min=1e-6) + + outputs = self.model(pixel_values=x) + semantic = outputs["semantic"] + label = outputs["label"] + + if self.semantic_postprocess == "none": + semantic_out = semantic + elif self.semantic_postprocess == "resize_logits": + semantic_out = F.interpolate( + semantic, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + elif self.semantic_postprocess == "argmax_lowres": + semantic_out = torch.argmax(semantic, dim=1).to(torch.uint8) + elif self.semantic_postprocess == "argmax_fullres": + semantic = F.interpolate( + semantic, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + semantic_out = torch.argmax(semantic, dim=1).to(torch.uint8) + else: + raise RuntimeError("semantic_postprocess inválido") + + if self.label_postprocess == "probs": + label_out = torch.softmax(label, dim=1) + else: + label_out = label + + return semantic_out, label_out + + +# ============================================================ +# Build/load +# ============================================================ + + +def build_visual_model( + backbone: str, + num_seg_classes: int, + num_label_classes: int, + device: torch.device, + input_hw: Tuple[int, int], +) -> VisualSegformerDualLabel: + H, W = input_hw + + base_model = SegformerForSemanticSegmentation.from_pretrained( + backbone, + num_labels=int(num_seg_classes), + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + base_model.config.output_hidden_states = True + base_model.to(device) + base_model.eval() + + with torch.no_grad(): + dummy = torch.zeros((1, 3, int(H), int(W)), dtype=torch.float32, device=device) + out = base_model(pixel_values=dummy) + logits = out.logits + feat = get_last_feat(out, logits) + feat_ch = int(feat.shape[1]) + + label_head = LabelHead( + feat_ch=feat_ch, + num_seg_classes=int(num_seg_classes), + num_label_classes=int(num_label_classes), + hidden=256, + dropout=0.2, + ).to(device) + label_head.eval() + + return VisualSegformerDualLabel(base_model=base_model, label_head=label_head).to(device) + + +def load_checkpoint_into_model(model: VisualSegformerDualLabel, checkpoint_path: Path): + ckpt = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + + if "model" not in ckpt: + raise RuntimeError( + f"Checkpoint não contém chave 'model': {checkpoint_path}. " + "Confirme se foi salvo pelo script de treino visual." + ) + + if "aux_head" not in ckpt: + raise RuntimeError( + f"Checkpoint não contém chave 'aux_head': {checkpoint_path}. " + "Este exportador é para o modelo visual dual_head_label." + ) + + model.base_model.load_state_dict(ckpt["model"], strict=True) + model.label_head.load_state_dict(ckpt["aux_head"], strict=True) + return ckpt + + +# ============================================================ +# Main +# ============================================================ + + +def main(): + parser = argparse.ArgumentParser() + + parser.add_argument("--config", default="config.json") + parser.add_argument("--checkpoint", default="") + parser.add_argument("--out", default="") + parser.add_argument("--labelmap", default="") + parser.add_argument("--norm_stats", default="") + + parser.add_argument("--opset", type=int, default=17) + parser.add_argument("--device", default="cuda", choices=["cuda", "cpu"]) + parser.add_argument("--dynamic-batch", action="store_true") + + parser.add_argument( + "--include-norm", + action="store_true", + help="Inclui normalização RGB x=(x-mean)/std dentro do ONNX. Runtime deve enviar RGB 0..1.", + ) + + parser.add_argument( + "--semantic-postprocess", + default="none", + choices=["none", "resize_logits", "argmax_lowres", "argmax_fullres"], + help="Define a saída semântica exportada.", + ) + + parser.add_argument( + "--label-postprocess", + default="logits", + choices=["logits", "probs"], + help="Define a saída da cabeça de status/corredor.", + ) + + 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) + + checkpoint_path, out_path, ckpt_name, mode, save_dir = resolve_default_paths( + args=args, + config=config, + config_dir=config_dir, + ) + + if mode != "label": + raise RuntimeError( + f"Este exportador foi preparado para dual_head_label. " + f"Modo detectado no config: {mode}" + ) + + if not checkpoint_path.is_file(): + raise FileNotFoundError(f"Checkpoint não encontrado: {checkpoint_path}") + + ensure_dir(out_path.parent) + + labelmap_path = resolve_labelmap_path(args, config, config_dir) + semantic_id2label, semantic_label2id, ignore_index = load_labelmap(labelmap_path) + num_seg_classes = len(semantic_id2label) + + # Lê ckpt antes para recuperar label_name_by_id, se necessário. + ckpt_for_meta = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + label_name_by_id, num_label_classes = resolve_label_classes(config, ckpt_for_meta) + + W, H = config.get("resolucao", [1024, 640]) + W = int(W) + H = int(H) + backbone = str(config.get("backbone", "nvidia/mit-b0")) + + norm_stats_path = None + norm_mean = None + norm_std = None + norm_channels = ["R", "G", "B"] + + if args.include_norm: + norm_stats_path = resolve_norm_stats_path(args, config, config_dir, save_dir) + if norm_stats_path is None: + raise RuntimeError("include_norm=True, mas norm_stats_path não foi resolvido.") + norm_mean, norm_std, norm_channels = load_rgb_norm_stats(norm_stats_path) + + 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 não disponível. Exportando em CPU.") + + print("==========================================") + print("Export Visual Worker SegFormer Dual Label para ONNX") + print(f"Config : {config_path}") + print(f"Checkpoint : {checkpoint_path}") + print(f"Output ONNX : {out_path}") + print(f"Save dir : {save_dir}") + print(f"Backbone : {backbone}") + print(f"Resolution : {W}x{H}") + print(f"Input shape : [1, 3, {H}, {W}]") + print(f"Semantic classes : {num_seg_classes} {semantic_id2label}") + print(f"Label classes : {num_label_classes} {label_name_by_id}") + print(f"Include norm : {args.include_norm}") + print(f"Norm stats : {norm_stats_path}") + print(f"Semantic output : {args.semantic_postprocess}") + print(f"Label output : {args.label_postprocess}") + print(f"Dynamic batch : {args.dynamic_batch}") + print(f"Opset : {args.opset}") + print(f"Device : {device}") + print("==========================================") + + print("[MODEL] Montando modelo visual...") + model = build_visual_model( + backbone=backbone, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + device=device, + input_hw=(H, W), + ) + + print("[CKPT] Carregando checkpoint...") + ckpt = load_checkpoint_into_model(model, checkpoint_path) + model.to(device) + model.eval() + + wrapper = VisualOnnxWrapper( + model=model, + include_norm=args.include_norm, + norm_mean=norm_mean, + norm_std=norm_std, + semantic_postprocess=args.semantic_postprocess, + label_postprocess=args.label_postprocess, + ).to(device) + wrapper.eval() + + dummy_input = torch.randn( + 1, + 3, + H, + W, + dtype=torch.float32, + device=device, + ) + + input_names = ["pixel_values"] + + if args.semantic_postprocess in ("argmax_lowres", "argmax_fullres"): + semantic_output_name = "semantic_mask" + else: + semantic_output_name = "semantic_logits" + + label_output_name = "label_probs" if args.label_postprocess == "probs" else "label_logits" + output_names = [semantic_output_name, label_output_name] + + dynamic_axes = None + if args.dynamic_batch: + dynamic_axes = { + "pixel_values": {0: "batch"}, + semantic_output_name: {0: "batch"}, + label_output_name: {0: "batch"}, + } + + print("[CHECK] Rodando forward PyTorch antes do export...") + with torch.no_grad(): + y = wrapper(dummy_input) + + for name, tensor in zip(output_names, y): + print(f" {name}: shape={tuple(tensor.shape)} dtype={tensor.dtype}") + + print("[EXPORT] Exportando ONNX...") + with torch.no_grad(): + torch.onnx.export( + wrapper, + dummy_input, + str(out_path), + export_params=True, + opset_version=int(args.opset), + do_constant_folding=True, + input_names=input_names, + output_names=output_names, + dynamic_axes=dynamic_axes, + ) + + print(f"[OK] ONNX salvo em: {out_path}") + + meta_path = out_path.with_suffix(".export_meta.json") + save_json(meta_path, { + "kind": "visual_worker_segformer_dual_label", + "config": str(config_path), + "checkpoint": str(checkpoint_path), + "onnx": str(out_path), + "ckpt_name": ckpt_name, + "mode": mode, + "backbone": backbone, + "input_shape": [1, 3, H, W], + "input_channel_names": ["R", "G", "B"], + "input_channel_indices": [0, 1, 2], + "semantic_output_name": semantic_output_name, + "label_output_name": label_output_name, + "output_names": output_names, + "semantic_postprocess": str(args.semantic_postprocess), + "label_postprocess": str(args.label_postprocess), + "opset": int(args.opset), + "dynamic_batch": bool(args.dynamic_batch), + "include_norm": bool(args.include_norm), + "norm_stats_path": str(norm_stats_path) if norm_stats_path is not None else None, + "norm_channels": norm_channels, + "norm_mean": norm_mean, + "norm_std": norm_std, + "semantic_id2label": semantic_id2label, + "semantic_label2id": semantic_label2id, + "ignore_index": int(ignore_index), + "label_name_by_id": label_name_by_id, + "num_label_classes": int(num_label_classes), + "checkpoint_epoch": ckpt.get("epoch", None), + "checkpoint_bests": ckpt.get("bests", None), + "checkpoint_extra": ckpt.get("extra", None), + }) + print(f"[OK] Metadata salvo em: {meta_path}") + + try: + import onnx + print("[ONNX] Verificando grafo com onnx.checker...") + onnx_model = onnx.load(str(out_path)) + onnx.checker.check_model(onnx_model) + print("[OK] onnx.checker passou.") + except ImportError: + print("[WARN] Pacote onnx não instalado. Pulei onnx.checker.") + print(" Instale com: pip install onnx") + except Exception as e: + print(f"[WARN] onnx.checker encontrou problema: {e}") + + print("\nExport finalizado.") + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/_11_validate_onnx.py b/Python/OAK/datasets/_11_validate_onnx.py new file mode 100644 index 000000000..5cf8fa308 --- /dev/null +++ b/Python/OAK/datasets/_11_validate_onnx.py @@ -0,0 +1,1518 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +_11_validate_visual_onnx.py + +Valida fidelidade entre: + - checkpoint PyTorch .pt do Visual Worker + - modelo ONNX .onnx exportado pelo _10_export_visual_onnx.py + +Contrato esperado: + - Entrada RGB [N, 3, H, W], float32 + - Cabeça semântica: semantic_logits ou semantic_mask + - Cabeça de status: label_logits ou label_probs + +Valida: + - Segmentação PyTorch vs ONNX: + logits/prob diff, argmax_equal, mIoU entre predições + - Label/status PyTorch vs ONNX: + logits/probs diff, top1_equal, classe prevista, confiança + +Exemplos: + +# ONNX exportado com --include-norm e saída semantic_logits + label_probs +python _11_validate_visual_onnx.py ^ + --config config.json ^ + --device cuda ^ + --onnx_provider cuda ^ + --torch_no_amp ^ + --onnx_has_norm ^ + --semantic_output_kind logits ^ + --label_output_kind probs ^ + --compare_at_input_size + +# TensorRT +python _11_validate_visual_onnx.py ^ + --config config.json ^ + --device cuda ^ + --onnx_provider tensorrt ^ + --torch_no_amp ^ + --onnx_has_norm ^ + --semantic_output_kind logits ^ + --label_output_kind probs ^ + --compare_at_input_size +""" + +from __future__ import annotations + +import os +import glob +import json +import time +import argparse +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import SegformerForSemanticSegmentation + + +# ============================================================ +# Utils básicos +# ============================================================ + + +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 or str(path_like).strip() == "": + return None + + p = Path(path_like) + if p.is_absolute(): + return p + + if base is None: + base = Path.cwd() + + return (base / p).resolve() + + +def _try_int(text: str) -> Optional[int]: + try: + return int(str(text).strip()) + except Exception: + return None + + +def clean_label_name(raw_name: str) -> str: + name = str(raw_name).strip() + + if "::" in name: + name = name.split("::", 1)[0].strip() + + if ":" in name: + left, right = name.split(":", 1) + right_clean = right.replace(",", "").replace(" ", "") + if right_clean.isdigit(): + name = left.strip() + + return name.strip() + + +def load_labelmap(labelmap_path: Path) -> Tuple[Dict[int, str], Dict[str, int], int]: + if not labelmap_path.is_file(): + raise FileNotFoundError(f"Labelmap não encontrado: {labelmap_path}") + + id2label: Dict[int, str] = {} + ignore_index = 255 + next_id = 0 + + with labelmap_path.open("r", encoding="utf-8") as f: + for raw_line in f: + line = raw_line.strip() + if not line or line.startswith("#"): + continue + + lower = line.lower() + if lower.startswith("ignore") or lower.startswith("ignore_index"): + for sep in ("=", ":", ",", " "): + if sep in line: + maybe = _try_int(line.split(sep)[-1]) + if maybe is not None: + ignore_index = maybe + break + continue + + cls_id: Optional[int] = None + cls_name: Optional[str] = None + + # Exemplo: navegavel:0,128,0:: + if "::" in line and ":" in line: + before = line.split("::", 1)[0].strip() + maybe_name = before.split(":", 1)[0].strip() + if maybe_name: + cls_id = next_id + cls_name = maybe_name + + if cls_id is None: + for sep in (":", ",", "\t", " "): + if sep in line: + parts = [p.strip() for p in line.split(sep) if p.strip()] + if len(parts) >= 2: + left_id = _try_int(parts[0]) + right_id = _try_int(parts[-1]) + + if left_id is not None: + cls_id = left_id + cls_name = sep.join(parts[1:]).strip() if sep in (":", ",") else " ".join(parts[1:]).strip() + break + + if right_id is not None: + cls_id = right_id + cls_name = sep.join(parts[:-1]).strip() if sep in (":", ",") else " ".join(parts[:-1]).strip() + break + + if cls_id is None: + cls_id = next_id + cls_name = line + + if cls_name is None or cls_name == "": + raise RuntimeError(f"Linha inválida no labelmap: {raw_line!r}") + + id2label[int(cls_id)] = clean_label_name(str(cls_name)) + next_id = max(next_id, int(cls_id) + 1) + + if not id2label: + raise RuntimeError(f"Labelmap vazio ou inválido: {labelmap_path}") + + ids_sorted = sorted(id2label.keys()) + if ids_sorted != list(range(len(ids_sorted))): + remap = {old_id: new_id for new_id, old_id in enumerate(ids_sorted)} + id2label = {remap[old_id]: name for old_id, name in id2label.items()} + + label2id = {name: idx for idx, name in id2label.items()} + return id2label, label2id, ignore_index + + +def get_visual_mode(config: dict) -> str: + use_mask2 = bool(config.get("dual_head_mask", config.get("dual_head", False))) + use_label = bool(config.get("dual_head_label", False)) + + if use_mask2 and use_label: + raise RuntimeError("Config inválido: dual_head_mask e dual_head_label ativos juntos.") + + if use_label: + return "label" + if use_mask2: + return "mask2" + return "single" + + +def get_save_suffix(mode: str) -> str: + if mode == "single": + return "_single" + if mode == "mask2": + return "_dual_mask" + if mode == "label": + return "_dual_label" + raise RuntimeError(f"Modo desconhecido: {mode}") + + +def resolve_labelmap_path(args, config: dict, config_dir: Path) -> Path: + explicit = resolve_path(args.labelmap, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + camera = str(config.get("camera", "oak-d")) + return (config_dir / camera / "dataset" / "labelmap.txt").resolve() + + +def resolve_default_paths(args, config: dict, config_dir: Path) -> Tuple[Path, Path, str, str, Path]: + camera = str(config.get("camera", "oak-d")) + model_key = str(config.get("modelo", "segformer_b0")) + model_name = str(config.get("model_name", "visual")) + mode = get_visual_mode(config) + suffix = get_save_suffix(mode) + + save_dir = config_dir / camera / "backup" / model_key / f"{model_name}{suffix}" + + if mode == "label": + default_ckpt_name = "best_label" + elif mode == "mask2": + default_ckpt_name = "best_mask2" + else: + default_ckpt_name = "best_main" + + ckpt_name = str(config.get("ckpt_test", default_ckpt_name)) + + checkpoint_path = resolve_path(args.checkpoint, Path.cwd()) + if checkpoint_path is None: + checkpoint_path = save_dir / f"{ckpt_name}.pt" + + onnx_path = resolve_path(args.onnx, Path.cwd()) + if onnx_path is None: + onnx_path = save_dir / f"{ckpt_name}.onnx" + + if 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 not onnx_path.is_file(): + raise FileNotFoundError( + f"ONNX não encontrado: {onnx_path}\n" + f"Dica: informe --onnx ou exporte antes com _10_export_visual_onnx.py." + ) + + return checkpoint_path.resolve(), onnx_path.resolve(), ckpt_name, mode, save_dir.resolve() + + +def resolve_norm_stats_path(args, config: dict, config_dir: Path, save_dir: Path) -> Optional[Path]: + explicit = resolve_path(args.norm_stats, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + p = save_dir / "norm_stats.json" + if p.is_file(): + return p.resolve() + + W, H = config.get("resolucao", [1024, 640]) + camera = str(config.get("camera", "oak-d")) + p = config_dir / camera / "dataset" / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json" + if p.is_file(): + return p.resolve() + + return p.resolve() + + +def load_rgb_norm_stats(path: Optional[Path]) -> Tuple[Optional[List[float]], Optional[List[float]], Optional[str], List[str]]: + if path is None or not path.is_file(): + if path is not None: + print(f"[NORM] norm_stats não encontrado: {path}") + print("[NORM] Sem norm_stats. Usando tensor 0..1 sem padronização.") + return None, None, None, ["R", "G", "B"] + + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + if names: + name_to_idx = {str(n).upper(): i for i, n in enumerate(names)} + required = ["R", "G", "B"] + missing = [ch for ch in required if ch not in name_to_idx] + if missing: + raise RuntimeError(f"norm_stats incompatível: faltam canais {missing}. channels={names}") + idx = [name_to_idx[ch] for ch in required] + mean_sel = [float(mean[i]) for i in idx] + std_sel = [float(std[i]) for i in idx] + names_sel = required + else: + if len(mean) < 3 or len(std) < 3: + raise RuntimeError(f"norm_stats precisa de pelo menos 3 valores RGB: {path}") + mean_sel = [float(mean[i]) for i in range(3)] + std_sel = [float(std[i]) for i in range(3)] + names_sel = ["R", "G", "B"] + + print(f"[NORM] usando {path}") + print(f"[NORM] channels={names_sel}") + print(f"[NORM] mean={mean_sel}") + print(f"[NORM] std ={std_sel}") + + return mean_sel, std_sel, str(path), names_sel + + +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 resolve_label_classes(config: dict, ckpt: Optional[dict] = None) -> Tuple[Dict[int, str], int]: + label_classes = config.get("label_classes", None) + + if label_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(label_classes)} + return label_name_by_id, len(label_name_by_id) + + if ckpt is not None: + extra = ckpt.get("extra", {}) if isinstance(ckpt, dict) else {} + maybe = extra.get("label_name_by_id", None) + if isinstance(maybe, dict) and maybe: + label_name_by_id = {int(k): str(v) for k, v in maybe.items()} + return label_name_by_id, max(label_name_by_id.keys()) + 1 + + maybe_config = extra.get("config", {}) if isinstance(extra, dict) else {} + maybe_classes = maybe_config.get("label_classes", None) if isinstance(maybe_config, dict) else None + if maybe_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(maybe_classes)} + return label_name_by_id, len(label_name_by_id) + + raise RuntimeError( + "Não consegui resolver label_classes. " + "Adicione config['label_classes'] ou use um checkpoint com extra['label_name_by_id']." + ) + + +# ============================================================ +# Dataset / imagens +# ============================================================ + + +@dataclass +class Sample: + img_path: str + mask_path: Optional[str] + mask2_path: Optional[str] + label_json_path: Optional[str] + label_npy_path: Optional[str] + group_name: str + filename: str + + +IMG_EXTS = (".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff", ".webp") + + +def find_by_stem(dir_path: str, filename: str, exts: Tuple[str, ...]) -> Optional[str]: + if not dir_path or not os.path.isdir(dir_path): + return None + + stem, _ = os.path.splitext(filename) + + for ext in exts: + p = os.path.join(dir_path, stem + ext) + if os.path.exists(p): + return p + + return None + + +def discover_samples(split_root: Path, max_samples: int = 20, start_idx: int = 0) -> List[Sample]: + """ + Descobre amostras no layout do visual worker: + + split/val/group//images/*.jpeg + split/val/group//masks/*.png + split/val/group//labels/*.json ou .npy + + O validate compara PyTorch vs ONNX usando a imagem. Máscara/label GT são + apenas metadados úteis no relatório, não entram na fidelidade PT vs ONNX. + """ + split_root = Path(split_root) + group_root = split_root / "group" + + if not group_root.is_dir(): + raise RuntimeError(f"Não achei pasta: {group_root}") + + samples: List[Sample] = [] + img_dirs = glob.glob(str(group_root / "**" / "images"), recursive=True) + img_dirs = [d for d in img_dirs if os.path.isdir(d)] + + for idir in sorted(img_dirs): + base = os.path.dirname(idir) + group_name = os.path.relpath(base, str(group_root)).replace("\\", "/") + mdir = os.path.join(base, "masks") + m2dir = os.path.join(base, "masks2") + ldir = os.path.join(base, "labels") + + img_paths: List[str] = [] + for ext in IMG_EXTS: + img_paths.extend(glob.glob(os.path.join(idir, f"*{ext}"))) + img_paths.extend(glob.glob(os.path.join(idir, f"*{ext.upper()}"))) + + for ip in sorted(set(img_paths)): + fn = os.path.basename(ip) + mask_path = find_by_stem(mdir, fn, (".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff")) + mask2_path = find_by_stem(m2dir, fn, (".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff")) + label_npy_path = find_by_stem(ldir, fn, (".npy",)) + label_json_path = find_by_stem(ldir, fn, (".json", ".txt")) + + samples.append(Sample( + img_path=ip, + mask_path=mask_path, + mask2_path=mask2_path, + label_json_path=label_json_path, + label_npy_path=label_npy_path, + group_name=group_name, + filename=fn, + )) + + if not samples: + raise RuntimeError(f"Nenhuma imagem encontrada em: {group_root}/**/images") + + start_idx = max(0, int(start_idx)) + selected = samples[start_idx:] + + if max_samples > 0: + selected = selected[:int(max_samples)] + + return selected + + +def discover_image_folder(folder: Path, max_samples: int = 20, start_idx: int = 0) -> List[Sample]: + folder = Path(folder) + samples: List[Sample] = [] + + img_paths: List[str] = [] + for ext in IMG_EXTS: + img_paths.extend(glob.glob(str(folder / f"*{ext}"))) + img_paths.extend(glob.glob(str(folder / f"*{ext.upper()}"))) + + for ip in sorted(set(img_paths)): + samples.append(Sample( + img_path=ip, + mask_path=None, + mask2_path=None, + label_json_path=None, + label_npy_path=None, + group_name="external", + filename=os.path.basename(ip), + )) + + if not samples: + raise RuntimeError(f"Nenhuma imagem encontrada em: {folder}") + + start_idx = max(0, int(start_idx)) + selected = samples[start_idx:] + + if max_samples > 0: + selected = selected[:int(max_samples)] + + return selected + + +def read_gt_label(sample: Sample) -> Tuple[Optional[int], Optional[str]]: + if sample.label_npy_path and os.path.exists(sample.label_npy_path): + try: + v = np.load(sample.label_npy_path) + return int(np.array(v).reshape(-1)[0]), None + except Exception: + pass + + if sample.label_json_path and os.path.exists(sample.label_json_path): + ext = os.path.splitext(sample.label_json_path)[1].lower() + + if ext == ".json": + with open(sample.label_json_path, "r", encoding="utf-8") as f: + data = json.load(f) + lid = data.get("label_id") + lname = data.get("estado_corredor") or data.get("label") or data.get("state") + return (int(lid) if lid is not None else None), lname + + if ext == ".txt": + with open(sample.label_json_path, "r", encoding="utf-8") as f: + txt = f.read().strip() + try: + return int(txt), None + except Exception: + return None, txt + + return None, None + + +def load_rgb_image(path: str | Path) -> np.ndarray: + img_bgr = cv2.imread(str(path), cv2.IMREAD_COLOR) + if img_bgr is None: + raise RuntimeError(f"Falha ao ler imagem: {path}") + + img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) + chw = np.transpose(img_rgb.astype(np.float32), (2, 0, 1)) / 255.0 + return np.clip(chw, 0.0, 1.0).astype(np.float32) + + +def resize_chw(chw: np.ndarray, target_hw: Tuple[int, int]) -> np.ndarray: + H, W = target_hw + if chw.shape[-2:] == (H, W): + return chw.astype(np.float32, copy=False) + hwc = np.transpose(chw, (1, 2, 0)) + hwc = cv2.resize(hwc, (W, H), interpolation=cv2.INTER_AREA) + return np.transpose(hwc, (2, 0, 1)).astype(np.float32) + + +# ============================================================ +# Métricas +# ============================================================ + + +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: + n, c, h, w = logits.shape + th, tw = target_hw + + if (h, w) == (th, tw): + return logits.astype(np.float32, copy=False) + + 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]]: + 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: + 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 mean_or_none(values: List[float]) -> Optional[float]: + return None if not values else float(np.mean(values)) + + +def min_or_none(values: List[float]) -> Optional[float]: + return None if not values else float(np.min(values)) + + +def max_or_none(values: List[float]) -> Optional[float]: + return None if not values else float(np.max(values)) + + +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 fmt_optional(v: Optional[float], casas: int = 8) -> str: + if v is None: + return "N/A" + return f"{float(v):.{casas}f}" + + +# ============================================================ +# Modelo PyTorch visual +# ============================================================ + + +class LabelHead(nn.Module): + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + hidden: int = 256, + dropout: float = 0.2, + ): + super().__init__() + in_ch = int(feat_ch) + int(num_seg_classes) + self.in_ch = in_ch + self.pool = nn.AdaptiveAvgPool2d((1, 1)) + self.net = nn.Sequential( + nn.Linear(in_ch, hidden), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + nn.Linear(hidden, num_label_classes), + ) + + def forward(self, feat: torch.Tensor, logits_seg: torch.Tensor) -> torch.Tensor: + feat = F.interpolate( + feat, + size=logits_seg.shape[-2:], + mode="bilinear", + align_corners=False, + ) + x = torch.cat([feat, logits_seg], dim=1) + x = self.pool(x).flatten(1) + return self.net(x) + + +def get_last_feat(out, logits: torch.Tensor) -> torch.Tensor: + if hasattr(out, "hidden_states") and out.hidden_states is not None: + feat = out.hidden_states[-1] + else: + feat = logits + + feat = F.interpolate( + feat, + size=logits.shape[-2:], + mode="bilinear", + align_corners=False, + ) + return feat + + +class VisualSegformerDualLabel(nn.Module): + def __init__(self, base_model: nn.Module, label_head: nn.Module): + super().__init__() + self.base_model = base_model + self.label_head = label_head + + def forward(self, pixel_values: torch.Tensor) -> Dict[str, torch.Tensor]: + out = self.base_model(pixel_values=pixel_values) + logits_seg = out.logits + feat = get_last_feat(out, logits_seg) + logits_label = self.label_head(feat, logits_seg) + return { + "semantic_logits": logits_seg, + "label_logits": logits_label, + "label_probs": torch.softmax(logits_label, dim=1), + } + + +def build_visual_model( + backbone: str, + num_seg_classes: int, + num_label_classes: int, + device: torch.device, + input_hw: Tuple[int, int], +) -> VisualSegformerDualLabel: + H, W = input_hw + + base_model = SegformerForSemanticSegmentation.from_pretrained( + backbone, + num_labels=int(num_seg_classes), + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + base_model.config.output_hidden_states = True + base_model.to(device) + base_model.eval() + + with torch.no_grad(): + dummy = torch.zeros((1, 3, int(H), int(W)), dtype=torch.float32, device=device) + out = base_model(pixel_values=dummy) + logits = out.logits + feat = get_last_feat(out, logits) + feat_ch = int(feat.shape[1]) + + label_head = LabelHead( + feat_ch=feat_ch, + num_seg_classes=int(num_seg_classes), + num_label_classes=int(num_label_classes), + hidden=256, + dropout=0.2, + ).to(device) + label_head.eval() + + return VisualSegformerDualLabel(base_model=base_model, label_head=label_head).to(device) + + +def load_checkpoint_into_model(model: VisualSegformerDualLabel, checkpoint_path: Path): + ckpt = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + + if "model" not in ckpt: + raise RuntimeError(f"Checkpoint não contém chave 'model': {checkpoint_path}") + + if "aux_head" not in ckpt: + raise RuntimeError(f"Checkpoint não contém chave 'aux_head': {checkpoint_path}") + + model.base_model.load_state_dict(ckpt["model"], strict=True) + model.label_head.load_state_dict(ckpt["aux_head"], strict=True) + return ckpt + + +class VisualTorchWrapper(nn.Module): + def __init__( + self, + model: VisualSegformerDualLabel, + resize_semantic_to_input: bool = False, + semantic_output_kind: str = "logits", + label_output_kind: str = "probs", + ): + super().__init__() + self.model = model + self.resize_semantic_to_input = bool(resize_semantic_to_input) + self.semantic_output_kind = str(semantic_output_kind).lower() + self.label_output_kind = str(label_output_kind).lower() + + if self.semantic_output_kind not in ("logits", "mask"): + raise RuntimeError(f"semantic_output_kind inválido: {self.semantic_output_kind}") + if self.label_output_kind not in ("logits", "probs"): + raise RuntimeError(f"label_output_kind inválido: {self.label_output_kind}") + + def forward(self, pixel_values: torch.Tensor) -> Dict[str, torch.Tensor]: + outputs = self.model(pixel_values=pixel_values) + semantic = outputs["semantic_logits"] + label_logits = outputs["label_logits"] + input_hw = pixel_values.shape[-2:] + + if self.resize_semantic_to_input or self.semantic_output_kind == "mask": + semantic = F.interpolate( + semantic, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + + if self.semantic_output_kind == "mask": + semantic_out = torch.argmax(semantic, dim=1).to(torch.uint8) + else: + semantic_out = semantic + + if self.label_output_kind == "probs": + label_out = torch.softmax(label_logits, dim=1) + else: + label_out = label_logits + + return { + "semantic": semantic_out, + "label": label_out, + } + + +# ============================================================ +# ONNX Runtime +# ============================================================ + + +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() + + 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")) + + 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), + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str(cache_dir), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(cache_dir), + "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}") + + if provider == "tensorrt" and "TensorrtExecutionProvider" not in active: + raise RuntimeError( + "TensorRTExecutionProvider foi solicitado, mas não ficou ativo. " + f"Providers ativos: {active}." + ) + + 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 get_onnx_output_pair(onnx_outputs: Dict[str, np.ndarray]) -> Tuple[np.ndarray, np.ndarray, str, str]: + keys = list(onnx_outputs.keys()) + + semantic_candidates = ["semantic_logits", "semantic_mask", "semantic", "output_semantic"] + label_candidates = ["label_probs", "label_logits", "label", "output_label"] + + semantic_name = None + label_name = None + + for c in semantic_candidates: + if c in onnx_outputs: + semantic_name = c + break + + for c in label_candidates: + if c in onnx_outputs: + label_name = c + break + + if semantic_name is None or label_name is None: + if len(keys) != 2: + raise RuntimeError(f"Esperava 2 outputs ONNX, recebi {keys}") + semantic_name = semantic_name or keys[0] + label_name = label_name or keys[1] + + return onnx_outputs[semantic_name], onnx_outputs[label_name], semantic_name, label_name + + +# ============================================================ +# 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("--labelmap", default="") + parser.add_argument("--norm_stats", default="") + + parser.add_argument("--split_folder", default="val", choices=["train", "val", "test"]) + parser.add_argument("--root_override", default=None) + parser.add_argument( + "--test_folder", + default=None, + help="Pasta externa com imagens soltas para validar fidelidade PT vs ONNX.", + ) + 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.", + ) + parser.add_argument("--trt_no_fp16", action="store_true") + parser.add_argument("--torch_no_amp", action="store_true") + + parser.add_argument( + "--onnx_has_norm", + action="store_true", + help="Use quando o ONNX foi exportado com --include-norm. Nesse caso o ONNX recebe RGB 0..1.", + ) + + parser.add_argument( + "--semantic_output_kind", + default="logits", + choices=["logits", "mask"], + help="Tipo da saída semântica do ONNX: logits ou mask.", + ) + + parser.add_argument( + "--label_output_kind", + default="probs", + choices=["logits", "probs"], + help="Tipo da saída label do ONNX: logits ou probs.", + ) + + parser.add_argument( + "--compare_at_input_size", + action="store_true", + help="Compara segmentação em HxW da entrada.", + ) + + 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) + + checkpoint_path, onnx_path, ckpt_name, mode, save_dir = resolve_default_paths(args, config, config_dir) + + if mode != "label": + raise RuntimeError(f"Este validate foi preparado para dual_head_label. Modo detectado: {mode}") + + labelmap_path = resolve_labelmap_path(args, config, config_dir) + semantic_id2label, semantic_label2id, ignore_index = load_labelmap(labelmap_path) + num_seg_classes = len(semantic_id2label) + + W, H = config.get("resolucao", [1024, 640]) + W = int(W) + H = int(H) + backbone = str(config.get("backbone", "nvidia/mit-b0")) + + ckpt_meta = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + label_name_by_id, num_label_classes = resolve_label_classes(config, ckpt_meta) + + norm_stats_path = resolve_norm_stats_path(args, config, config_dir, save_dir) + mean, std, norm_stats_used, norm_channels = load_rgb_norm_stats(norm_stats_path) + + if args.test_folder: + root = resolve_path(args.test_folder, Path.cwd()) + if root is None or not root.is_dir(): + raise FileNotFoundError(f"Pasta de teste não encontrada: {root}") + samples = discover_image_folder( + folder=root, + max_samples=args.max_samples, + start_idx=args.start_idx, + ) + else: + if args.root_override: + root = resolve_path(args.root_override, Path.cwd()) + else: + camera = str(config.get("camera", "oak-d")) + root = (config_dir / camera / "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 = discover_samples( + split_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 Visual Worker PyTorch vs ONNX") + print(f"Config : {config_path}") + print(f"Checkpoint : {checkpoint_path}") + print(f"ONNX : {onnx_path}") + print(f"Root : {root}") + print(f"Samples : {len(samples)}") + print(f"Backbone : {backbone}") + print(f"Input shape : [1, 3, {H}, {W}]") + print(f"Semantic classes : {num_seg_classes} {semantic_id2label}") + print(f"Label classes : {num_label_classes} {label_name_by_id}") + 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 has norm : {args.onnx_has_norm}") + print(f"Semantic output kind: {args.semantic_output_kind}") + print(f"Label output kind : {args.label_output_kind}") + print(f"Compare HxW : {args.compare_at_input_size}") + print("==========================================") + + print("[MODEL] Montando PyTorch...") + model = build_visual_model( + backbone=backbone, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + device=device, + input_hw=(H, W), + ) + + print("[CKPT] Carregando checkpoint...") + ckpt = load_checkpoint_into_model(model, checkpoint_path) + model.to(device) + model.eval() + + torch_wrapper = VisualTorchWrapper( + model=model, + resize_semantic_to_input=args.compare_at_input_size or args.semantic_output_kind == "mask", + semantic_output_kind=args.semantic_output_kind, + label_output_kind=args.label_output_kind, + ).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_name = onnx_session.get_inputs()[0].name + onnx_output_names = [o.name for o in onnx_session.get_outputs()] + print(f"[ONNX] input name: {onnx_input_name}") + print(f"[ONNX] outputs: {onnx_output_names}") + + semantic_acc = { + "n": 0, + "logits_abs_mean": [], + "logits_abs_max": [], + "prob_abs_mean": [], + "prob_abs_max": [], + "argmax_equal_ratio": [], + "pred_miou": [], + "pred_iou_per_class": [], + } + + label_acc = { + "n": 0, + "abs_mean": [], + "abs_max": [], + "top1_equal": [], + "top1_pt": [], + "top1_onnx": [], + "conf_pt": [], + "conf_onnx": [], + } + + sample_reports = [] + + for i, sample in enumerate(samples): + chw01 = load_rgb_image(sample.img_path) + chw01 = resize_chw(chw01, target_hw=(H, W)) + chw_norm = normalize_numpy_chw(chw01, mean=mean, std=std) + + # PyTorch puro espera tensor já normalizado. + x_torch_np = np.expand_dims(chw_norm, axis=0).astype(np.float32) + + # ONNX com include_norm recebe RGB cru 0..1. + if args.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_semantic = torch_outputs_t["semantic"].detach().float().cpu().numpy() + torch_label = torch_outputs_t["label"].detach().float().cpu().numpy() + + t0 = time.perf_counter() + onnx_raw = run_onnx(onnx_session, onnx_input_name, x_onnx_np) + onnx_ms = (time.perf_counter() - t0) * 1000.0 + + onnx_semantic, onnx_label, onnx_semantic_name, onnx_label_name = get_onnx_output_pair(onnx_raw) + + gt_label_id, gt_label_name = read_gt_label(sample) + + report_item = { + "idx": i, + "image": str(sample.img_path), + "mask": str(sample.mask_path) if sample.mask_path else None, + "label_json": str(sample.label_json_path) if sample.label_json_path else None, + "label_npy": str(sample.label_npy_path) if sample.label_npy_path else None, + "group_name": sample.group_name, + "filename": sample.filename, + "gt_label_id": gt_label_id, + "gt_label_name": gt_label_name, + "torch_ms": float(torch_ms), + "onnx_ms": float(onnx_ms), + "onnx_semantic_name": onnx_semantic_name, + "onnx_label_name": onnx_label_name, + } + + print(f"[{i + 1:03d}/{len(samples):03d}] {sample.group_name}/{sample.filename} | torch={torch_ms:.2f}ms | onnx={onnx_ms:.2f}ms") + + # ==================================================== + # Segmentação + # ==================================================== + if args.semantic_output_kind == "mask": + pt_mask = np.asarray(torch_semantic).astype(np.uint8) + ox_mask = np.asarray(onnx_semantic).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)) + iou_per_class, miou, present_classes = compute_mask_iou_between_preds( + pt_mask, + ox_mask, + num_classes=num_seg_classes, + ) + + semantic_report = { + "torch_shape": list(np.asarray(torch_semantic).shape), + "onnx_shape": list(np.asarray(onnx_semantic).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": equal_ratio, + "pred_miou_torch_vs_onnx": miou, + "pred_iou_per_class": [None if x is None else float(x) for x in iou_per_class], + "present_classes": [int(x) for x in present_classes], + } + + print( + f" semantic " + f"shape PT={tuple(np.asarray(torch_semantic).shape)} ONNX={tuple(np.asarray(onnx_semantic).shape)} " + f"CMP={tuple(pt_mask.shape)} | " + f"mask_equal={equal_ratio * 100:.4f}% " + f"mIoU={miou:.8f}" + ) + + else: + pt = np.asarray(torch_semantic).astype(np.float32) + ox = np.asarray(onnx_semantic).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 semântico incompatível: 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)) + iou_per_class, miou, present_classes = compute_mask_iou_between_preds( + pred_pt, + pred_ox, + num_classes=num_seg_classes, + ) + + semantic_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": equal_ratio, + "pred_miou_torch_vs_onnx": miou, + "pred_iou_per_class": [None if x is None else float(x) for x in iou_per_class], + "present_classes": [int(x) for x in present_classes], + } + + print( + f" semantic " + f"shape PT={tuple(pt.shape)} ONNX={tuple(ox.shape)} CMP={tuple(pt_cmp.shape)} | " + f"logit_mean={semantic_report['logits_abs_mean']:.6g} " + f"prob_mean={semantic_report['prob_abs_mean']:.6g} " + f"argmax_equal={equal_ratio * 100:.4f}% " + f"mIoU={miou:.8f}" + ) + + semantic_acc["n"] += 1 + if semantic_report["logits_abs_mean"] is not None: + semantic_acc["logits_abs_mean"].append(semantic_report["logits_abs_mean"]) + semantic_acc["logits_abs_max"].append(semantic_report["logits_abs_max"]) + semantic_acc["prob_abs_mean"].append(semantic_report["prob_abs_mean"]) + semantic_acc["prob_abs_max"].append(semantic_report["prob_abs_max"]) + semantic_acc["argmax_equal_ratio"].append(semantic_report["argmax_equal_ratio"]) + semantic_acc["pred_miou"].append(semantic_report["pred_miou_torch_vs_onnx"]) + semantic_acc["pred_iou_per_class"].append(semantic_report["pred_iou_per_class"]) + + # ==================================================== + # Label/status + # ==================================================== + pt_label = np.asarray(torch_label).astype(np.float32) + ox_label = np.asarray(onnx_label).astype(np.float32) + + if pt_label.shape != ox_label.shape: + raise RuntimeError(f"Shape label incompatível: torch={pt_label.shape}, onnx={ox_label.shape}") + + if args.label_output_kind == "logits": + pt_label_probs = softmax_np(pt_label, axis=1) + ox_label_probs = softmax_np(ox_label, axis=1) + label_diff_base = np.abs(pt_label - ox_label) + else: + pt_label_probs = pt_label + ox_label_probs = ox_label + label_diff_base = np.abs(pt_label_probs - ox_label_probs) + + top1_pt = int(np.argmax(pt_label_probs, axis=1)[0]) + top1_ox = int(np.argmax(ox_label_probs, axis=1)[0]) + conf_pt = float(pt_label_probs[0, top1_pt]) + conf_ox = float(ox_label_probs[0, top1_ox]) + top1_equal = bool(top1_pt == top1_ox) + + label_report = { + "torch_shape": list(pt_label.shape), + "onnx_shape": list(ox_label.shape), + "abs_mean": float(label_diff_base.mean()), + "abs_max": float(label_diff_base.max()), + "top1_equal": top1_equal, + "top1_torch": top1_pt, + "top1_onnx": top1_ox, + "top1_torch_name": label_name_by_id.get(top1_pt, str(top1_pt)), + "top1_onnx_name": label_name_by_id.get(top1_ox, str(top1_ox)), + "conf_torch": conf_pt, + "conf_onnx": conf_ox, + "gt_label_id": gt_label_id, + "gt_label_name": gt_label_name, + "probs_torch": pt_label_probs.reshape(-1).astype(float).tolist(), + "probs_onnx": ox_label_probs.reshape(-1).astype(float).tolist(), + } + + label_acc["n"] += 1 + label_acc["abs_mean"].append(label_report["abs_mean"]) + label_acc["abs_max"].append(label_report["abs_max"]) + label_acc["top1_equal"].append(1.0 if top1_equal else 0.0) + label_acc["top1_pt"].append(top1_pt) + label_acc["top1_onnx"].append(top1_ox) + label_acc["conf_pt"].append(conf_pt) + label_acc["conf_onnx"].append(conf_ox) + + print( + f" label " + f"shape PT={tuple(pt_label.shape)} ONNX={tuple(ox_label.shape)} | " + f"abs_mean={label_report['abs_mean']:.6g} " + f"abs_max={label_report['abs_max']:.6g} " + f"top1_equal={top1_equal} " + f"PT={top1_pt}:{label_report['top1_torch_name']}({conf_pt:.4f}) " + f"ONNX={top1_ox}:{label_report['top1_onnx_name']}({conf_ox:.4f}) " + f"GT={gt_label_id}:{gt_label_name}" + ) + + report_item["semantic"] = semantic_report + report_item["label"] = label_report + sample_reports.append(report_item) + + # ======================================================== + # Resumo + # ======================================================== + ious_raw = semantic_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) + + semantic_summary = { + "n": int(semantic_acc["n"]), + "logits_abs_mean_avg": mean_or_none(semantic_acc["logits_abs_mean"]), + "logits_abs_mean_max": max_or_none(semantic_acc["logits_abs_mean"]), + "logits_abs_max_avg": mean_or_none(semantic_acc["logits_abs_max"]), + "logits_abs_max_max": max_or_none(semantic_acc["logits_abs_max"]), + "prob_abs_mean_avg": mean_or_none(semantic_acc["prob_abs_mean"]), + "prob_abs_mean_max": max_or_none(semantic_acc["prob_abs_mean"]), + "prob_abs_max_avg": mean_or_none(semantic_acc["prob_abs_max"]), + "prob_abs_max_max": max_or_none(semantic_acc["prob_abs_max"]), + "argmax_equal_ratio_avg": float(np.mean(semantic_acc["argmax_equal_ratio"])), + "argmax_equal_ratio_min": float(np.min(semantic_acc["argmax_equal_ratio"])), + "pred_miou_avg": float(np.mean(semantic_acc["pred_miou"])), + "pred_miou_min": float(np.min(semantic_acc["pred_miou"])), + "pred_iou_per_class_avg": nanmean_list(ious_arr), + "pred_iou_per_class_min": nanmin_list(ious_arr), + } + + label_summary = { + "n": int(label_acc["n"]), + "abs_mean_avg": mean_or_none(label_acc["abs_mean"]), + "abs_mean_max": max_or_none(label_acc["abs_mean"]), + "abs_max_avg": mean_or_none(label_acc["abs_max"]), + "abs_max_max": max_or_none(label_acc["abs_max"]), + "top1_equal_ratio_avg": float(np.mean(label_acc["top1_equal"])), + "top1_equal_ratio_min": float(np.min(label_acc["top1_equal"])), + "conf_torch_avg": mean_or_none(label_acc["conf_pt"]), + "conf_onnx_avg": mean_or_none(label_acc["conf_onnx"]), + } + + summary = { + "kind": "visual_worker_validate_onnx", + "config": str(config_path), + "checkpoint": str(checkpoint_path), + "ckpt_name": ckpt_name, + "onnx": str(onnx_path), + "root": str(root), + "samples": len(samples), + "input_shape": [1, 3, H, W], + "input_channel_names": ["R", "G", "B"], + "semantic_id2label": semantic_id2label, + "label_name_by_id": label_name_by_id, + "norm_stats_used": norm_stats_used, + "norm_channels": norm_channels, + "torch_amp": bool(not args.torch_no_amp and device.type == "cuda"), + "onnx_provider": args.onnx_provider, + "onnx_has_norm": bool(args.onnx_has_norm), + "semantic_output_kind": args.semantic_output_kind, + "label_output_kind": args.label_output_kind, + "compare_at_input_size": bool(args.compare_at_input_size), + "trt_home": args.trt_home or os.environ.get("TRT_HOME", None), + "trt_fp16": bool(not args.trt_no_fp16), + "semantic": semantic_summary, + "label": label_summary, + "sample_reports": sample_reports, + } + + print("\n========== RESUMO ==========") + + print("\n[semantic]") + print(f" logits_abs_mean avg : {fmt_optional(semantic_summary['logits_abs_mean_avg'])}") + print(f" prob_abs_mean avg : {fmt_optional(semantic_summary['prob_abs_mean_avg'])}") + print(f" argmax_equal avg : {semantic_summary['argmax_equal_ratio_avg'] * 100:.4f}%") + print(f" argmax_equal min : {semantic_summary['argmax_equal_ratio_min'] * 100:.4f}%") + print(f" pred_mIoU avg : {semantic_summary['pred_miou_avg']:.8f}") + print(f" pred_mIoU min : {semantic_summary['pred_miou_min']:.8f}") + print( + " IoU/classes avg : " + f"{[None if x is None else round(float(x), 6) for x in semantic_summary['pred_iou_per_class_avg']]}" + ) + print( + " IoU/classes min : " + f"{[None if x is None else round(float(x), 6) for x in semantic_summary['pred_iou_per_class_min']]}" + ) + + print("\n[label]") + print(f" abs_mean avg : {fmt_optional(label_summary['abs_mean_avg'])}") + print(f" abs_max max : {fmt_optional(label_summary['abs_max_max'])}") + print(f" top1_equal avg : {label_summary['top1_equal_ratio_avg'] * 100:.4f}%") + print(f" top1_equal min : {label_summary['top1_equal_ratio_min'] * 100:.4f}%") + print(f" conf_torch avg : {fmt_optional(label_summary['conf_torch_avg'])}") + print(f" conf_onnx avg : {fmt_optional(label_summary['conf_onnx_avg'])}") + + if args.save_report: + report_path = resolve_path(args.save_report, Path.cwd()) + 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() diff --git a/Python/OAK/datasets/_12_benchmark_onnx.py b/Python/OAK/datasets/_12_benchmark_onnx.py new file mode 100644 index 000000000..8cda3c8f0 --- /dev/null +++ b/Python/OAK/datasets/_12_benchmark_onnx.py @@ -0,0 +1,1249 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +""" +_12_benchmark_visual_onnx.py + +Benchmark PyTorch vs ONNX Runtime para o Visual Worker. + +Contrato esperado: + - Modelo: SegFormer + cabeça auxiliar LabelHead + - Entrada: RGB [N, 3, H, W], float32 + - ONNX exportado pelo _10_export_visual_onnx.py + - Saídas ONNX típicas: + semantic_logits: [N, 2, H, W] + label_probs : [N, 8] + +Mede: + - PyTorch FP32 + - PyTorch AMP/FP16 + - ONNX Runtime CUDA / CPU / TensorRT + +Importante: + - Este benchmark mede só inferência do modelo. + - As imagens são carregadas e pré-processadas antes da medição. + - Para ONNX com --include-norm, use --onnx_has_norm. + +Exemplos: + +# Benchmark completo CUDA +python _12_benchmark_visual_onnx.py ^ + --config config.json ^ + --onnx_provider cuda ^ + --onnx_has_norm ^ + --max_samples 50 ^ + --warmup 10 ^ + --repeat 5 + +# Benchmark TensorRT, medindo só ONNX TensorRT +python _12_benchmark_visual_onnx.py ^ + --config config.json ^ + --onnx_provider tensorrt ^ + --onnx_has_norm ^ + --max_samples 50 ^ + --warmup 20 ^ + --repeat 10 ^ + --skip_torch_fp32 ^ + --skip_torch_amp + +# Pasta externa com imagens soltas +python _12_benchmark_visual_onnx.py ^ + --config config.json ^ + --test_folder .\oak-d\dataset\split\train\group\naonavegavel_navegavel\images\ ^ + --onnx_provider tensorrt ^ + --onnx_has_norm ^ + --warmup 20 ^ + --repeat 10 +""" + +from __future__ import annotations + +import os +import gc +import csv +import glob +import json +import time +import argparse +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, List, Optional, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import SegformerForSemanticSegmentation + + +# ============================================================ +# Utils gerais +# ============================================================ + + +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 save_csv(path: str | Path, rows: List[dict]): + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + + if not rows: + return + + keys = list(rows[0].keys()) + with path.open("w", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=keys) + w.writeheader() + w.writerows(rows) + + +def resolve_path(path_like: Optional[str], base: Optional[Path] = None) -> Optional[Path]: + if path_like is None or str(path_like).strip() == "": + return None + + p = Path(path_like) + if p.is_absolute(): + return p + + if base is None: + base = Path.cwd() + + return (base / p).resolve() + + +def synchronize_if_cuda(device: torch.device): + if device.type == "cuda": + torch.cuda.synchronize() + + +def clear_cuda(): + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + + +def summarize_times(times_ms: List[float]) -> dict: + arr = np.asarray(times_ms, dtype=np.float64) + + if arr.size == 0: + return { + "n": 0, + "mean_ms": 0.0, + "median_ms": 0.0, + "min_ms": 0.0, + "max_ms": 0.0, + "p90_ms": 0.0, + "p95_ms": 0.0, + "p99_ms": 0.0, + "fps_mean": 0.0, + "fps_median": 0.0, + "fps_p95_latency": 0.0, + "fps_p99_latency": 0.0, + } + + mean_ms = float(arr.mean()) + median_ms = float(np.median(arr)) + p95_ms = float(np.percentile(arr, 95)) + p99_ms = float(np.percentile(arr, 99)) + + return { + "n": int(arr.size), + "mean_ms": mean_ms, + "median_ms": median_ms, + "min_ms": float(arr.min()), + "max_ms": float(arr.max()), + "p90_ms": float(np.percentile(arr, 90)), + "p95_ms": p95_ms, + "p99_ms": p99_ms, + "fps_mean": float(1000.0 / mean_ms) if mean_ms > 0 else 0.0, + "fps_median": float(1000.0 / median_ms) if median_ms > 0 else 0.0, + "fps_p95_latency": float(1000.0 / p95_ms) if p95_ms > 0 else 0.0, + "fps_p99_latency": float(1000.0 / p99_ms) if p99_ms > 0 else 0.0, + } + + +# ============================================================ +# Labelmap / paths / config +# ============================================================ + + +def _try_int(text: str) -> Optional[int]: + try: + return int(str(text).strip()) + except Exception: + return None + + +def clean_label_name(raw_name: str) -> str: + name = str(raw_name).strip() + + if "::" in name: + name = name.split("::", 1)[0].strip() + + if ":" in name: + left, right = name.split(":", 1) + right_clean = right.replace(",", "").replace(" ", "") + if right_clean.isdigit(): + name = left.strip() + + return name.strip() + + +def load_labelmap(labelmap_path: Path) -> Tuple[Dict[int, str], Dict[str, int], int]: + if not labelmap_path.is_file(): + raise FileNotFoundError(f"Labelmap não encontrado: {labelmap_path}") + + id2label: Dict[int, str] = {} + ignore_index = 255 + next_id = 0 + + with labelmap_path.open("r", encoding="utf-8") as f: + for raw_line in f: + line = raw_line.strip() + if not line or line.startswith("#"): + continue + + lower = line.lower() + if lower.startswith("ignore") or lower.startswith("ignore_index"): + for sep in ("=", ":", ",", " "): + if sep in line: + maybe = _try_int(line.split(sep)[-1]) + if maybe is not None: + ignore_index = maybe + break + continue + + cls_id: Optional[int] = None + cls_name: Optional[str] = None + + if "::" in line and ":" in line: + before = line.split("::", 1)[0].strip() + maybe_name = before.split(":", 1)[0].strip() + if maybe_name: + cls_id = next_id + cls_name = maybe_name + + if cls_id is None: + for sep in (":", ",", "\t", " "): + if sep in line: + parts = [p.strip() for p in line.split(sep) if p.strip()] + if len(parts) >= 2: + left_id = _try_int(parts[0]) + right_id = _try_int(parts[-1]) + + if left_id is not None: + cls_id = left_id + cls_name = sep.join(parts[1:]).strip() if sep in (":", ",") else " ".join(parts[1:]).strip() + break + + if right_id is not None: + cls_id = right_id + cls_name = sep.join(parts[:-1]).strip() if sep in (":", ",") else " ".join(parts[:-1]).strip() + break + + if cls_id is None: + cls_id = next_id + cls_name = line + + if cls_name is None or cls_name == "": + raise RuntimeError(f"Linha inválida no labelmap: {raw_line!r}") + + id2label[int(cls_id)] = clean_label_name(str(cls_name)) + next_id = max(next_id, int(cls_id) + 1) + + if not id2label: + raise RuntimeError(f"Labelmap vazio ou inválido: {labelmap_path}") + + ids_sorted = sorted(id2label.keys()) + if ids_sorted != list(range(len(ids_sorted))): + remap = {old_id: new_id for new_id, old_id in enumerate(ids_sorted)} + id2label = {remap[old_id]: name for old_id, name in id2label.items()} + + label2id = {name: idx for idx, name in id2label.items()} + return id2label, label2id, ignore_index + + +def get_visual_mode(config: dict) -> str: + use_mask2 = bool(config.get("dual_head_mask", config.get("dual_head", False))) + use_label = bool(config.get("dual_head_label", False)) + + if use_mask2 and use_label: + raise RuntimeError("Config inválido: dual_head_mask e dual_head_label ativos juntos.") + + if use_label: + return "label" + if use_mask2: + return "mask2" + return "single" + + +def get_save_suffix(mode: str) -> str: + if mode == "single": + return "_single" + if mode == "mask2": + return "_dual_mask" + if mode == "label": + return "_dual_label" + raise RuntimeError(f"Modo desconhecido: {mode}") + + +def resolve_labelmap_path(args, config: dict, config_dir: Path) -> Path: + explicit = resolve_path(args.labelmap, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + camera = str(config.get("camera", "oak-d")) + return (config_dir / camera / "dataset" / "labelmap.txt").resolve() + + +def resolve_default_paths(args, config: dict, config_dir: Path) -> Tuple[Path, Path, str, str, Path]: + camera = str(config.get("camera", "oak-d")) + model_key = str(config.get("modelo", "segformer_b0")) + model_name = str(config.get("model_name", "visual")) + mode = get_visual_mode(config) + suffix = get_save_suffix(mode) + + save_dir = config_dir / camera / "backup" / model_key / f"{model_name}{suffix}" + + if mode == "label": + default_ckpt_name = "best_label" + elif mode == "mask2": + default_ckpt_name = "best_mask2" + else: + default_ckpt_name = "best_main" + + ckpt_name = str(config.get("ckpt_test", default_ckpt_name)) + + checkpoint_path = resolve_path(args.checkpoint, Path.cwd()) + if checkpoint_path is None: + checkpoint_path = save_dir / f"{ckpt_name}.pt" + + onnx_path = resolve_path(args.onnx, Path.cwd()) + if onnx_path is None: + onnx_path = save_dir / f"{ckpt_name}.onnx" + + if 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 not onnx_path.is_file(): + raise FileNotFoundError( + f"ONNX não encontrado: {onnx_path}\n" + f"Dica: informe --onnx ou exporte antes com _10_export_visual_onnx.py." + ) + + return checkpoint_path.resolve(), onnx_path.resolve(), ckpt_name, mode, save_dir.resolve() + + +def resolve_norm_stats_path(args, config: dict, config_dir: Path, save_dir: Path) -> Optional[Path]: + explicit = resolve_path(args.norm_stats, Path.cwd()) + if explicit is not None: + return explicit.resolve() + + p = save_dir / "norm_stats.json" + if p.is_file(): + return p.resolve() + + W, H = config.get("resolucao", [1024, 640]) + camera = str(config.get("camera", "oak-d")) + p = config_dir / camera / "dataset" / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json" + if p.is_file(): + return p.resolve() + + return p.resolve() + + +def load_rgb_norm_stats(path: Optional[Path]) -> Tuple[Optional[List[float]], Optional[List[float]], Optional[str], List[str]]: + if path is None or not path.is_file(): + if path is not None: + print(f"[NORM] norm_stats não encontrado: {path}") + print("[NORM] Sem norm_stats. Usando tensor 0..1 sem padronização.") + return None, None, None, ["R", "G", "B"] + + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + if names: + name_to_idx = {str(n).upper(): i for i, n in enumerate(names)} + required = ["R", "G", "B"] + missing = [ch for ch in required if ch not in name_to_idx] + if missing: + raise RuntimeError(f"norm_stats incompatível: faltam canais {missing}. channels={names}") + idx = [name_to_idx[ch] for ch in required] + mean_sel = [float(mean[i]) for i in idx] + std_sel = [float(std[i]) for i in idx] + names_sel = required + else: + if len(mean) < 3 or len(std) < 3: + raise RuntimeError(f"norm_stats precisa de pelo menos 3 valores RGB: {path}") + mean_sel = [float(mean[i]) for i in range(3)] + std_sel = [float(std[i]) for i in range(3)] + names_sel = ["R", "G", "B"] + + print(f"[NORM] usando {path}") + print(f"[NORM] channels={names_sel}") + print(f"[NORM] mean={mean_sel}") + print(f"[NORM] std ={std_sel}") + + return mean_sel, std_sel, str(path), names_sel + + +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 resolve_label_classes(config: dict, ckpt: Optional[dict] = None) -> Tuple[Dict[int, str], int]: + label_classes = config.get("label_classes", None) + + if label_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(label_classes)} + return label_name_by_id, len(label_name_by_id) + + if ckpt is not None: + extra = ckpt.get("extra", {}) if isinstance(ckpt, dict) else {} + maybe = extra.get("label_name_by_id", None) + if isinstance(maybe, dict) and maybe: + label_name_by_id = {int(k): str(v) for k, v in maybe.items()} + return label_name_by_id, max(label_name_by_id.keys()) + 1 + + maybe_config = extra.get("config", {}) if isinstance(extra, dict) else {} + maybe_classes = maybe_config.get("label_classes", None) if isinstance(maybe_config, dict) else None + if maybe_classes is not None: + label_name_by_id = {i: str(name) for i, name in enumerate(maybe_classes)} + return label_name_by_id, len(label_name_by_id) + + raise RuntimeError( + "Não consegui resolver label_classes. " + "Adicione config['label_classes'] ou use um checkpoint com extra['label_name_by_id']." + ) + + +# ============================================================ +# Dataset / imagens +# ============================================================ + + +@dataclass +class Sample: + img_path: str + group_name: str + filename: str + + +IMG_EXTS = (".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff", ".webp") + + +def discover_samples(split_root: Path, max_samples: int = 50, start_idx: int = 0) -> List[Sample]: + split_root = Path(split_root) + group_root = split_root / "group" + + if not group_root.is_dir(): + raise RuntimeError(f"Não achei pasta: {group_root}") + + samples: List[Sample] = [] + img_dirs = glob.glob(str(group_root / "**" / "images"), recursive=True) + img_dirs = [d for d in img_dirs if os.path.isdir(d)] + + for idir in sorted(img_dirs): + base = os.path.dirname(idir) + group_name = os.path.relpath(base, str(group_root)).replace("\\", "/") + + img_paths: List[str] = [] + for ext in IMG_EXTS: + img_paths.extend(glob.glob(os.path.join(idir, f"*{ext}"))) + img_paths.extend(glob.glob(os.path.join(idir, f"*{ext.upper()}"))) + + for ip in sorted(set(img_paths)): + samples.append(Sample( + img_path=ip, + group_name=group_name, + filename=os.path.basename(ip), + )) + + if not samples: + raise RuntimeError(f"Nenhuma imagem encontrada em: {group_root}/**/images") + + start_idx = max(0, int(start_idx)) + selected = samples[start_idx:] + + if max_samples > 0: + selected = selected[:int(max_samples)] + + return selected + + +def discover_image_folder(folder: Path, max_samples: int = 50, start_idx: int = 0) -> List[Sample]: + folder = Path(folder) + samples: List[Sample] = [] + + img_paths: List[str] = [] + for ext in IMG_EXTS: + img_paths.extend(glob.glob(str(folder / f"*{ext}"))) + img_paths.extend(glob.glob(str(folder / f"*{ext.upper()}"))) + + for ip in sorted(set(img_paths)): + samples.append(Sample( + img_path=ip, + group_name="external", + filename=os.path.basename(ip), + )) + + if not samples: + raise RuntimeError(f"Nenhuma imagem encontrada em: {folder}") + + start_idx = max(0, int(start_idx)) + selected = samples[start_idx:] + + if max_samples > 0: + selected = selected[:int(max_samples)] + + return selected + + +def load_rgb_image(path: str | Path) -> np.ndarray: + img_bgr = cv2.imread(str(path), cv2.IMREAD_COLOR) + if img_bgr is None: + raise RuntimeError(f"Falha ao ler imagem: {path}") + + img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) + chw = np.transpose(img_rgb.astype(np.float32), (2, 0, 1)) / 255.0 + return np.clip(chw, 0.0, 1.0).astype(np.float32) + + +def resize_chw(chw: np.ndarray, target_hw: Tuple[int, int]) -> np.ndarray: + H, W = target_hw + if chw.shape[-2:] == (H, W): + return chw.astype(np.float32, copy=False) + + hwc = np.transpose(chw, (1, 2, 0)) + hwc = cv2.resize(hwc, (W, H), interpolation=cv2.INTER_AREA) + return np.transpose(hwc, (2, 0, 1)).astype(np.float32) + + +def load_inputs_as_numpy( + samples: List[Sample], + mean: Optional[List[float]], + std: Optional[List[float]], + target_hw: Tuple[int, int], + normalize_input: bool = True, +) -> List[np.ndarray]: + xs: List[np.ndarray] = [] + + for s in samples: + chw01 = load_rgb_image(s.img_path) + chw01 = resize_chw(chw01, target_hw=target_hw) + + if normalize_input: + chw = normalize_numpy_chw(chw01, mean=mean, std=std) + else: + chw = chw01.astype(np.float32, copy=False) + + xs.append(np.expand_dims(chw, axis=0).astype(np.float32)) + + return xs + + +# ============================================================ +# Modelo PyTorch visual +# ============================================================ + + +class LabelHead(nn.Module): + def __init__( + self, + feat_ch: int, + num_seg_classes: int, + num_label_classes: int, + hidden: int = 256, + dropout: float = 0.2, + ): + super().__init__() + in_ch = int(feat_ch) + int(num_seg_classes) + self.in_ch = in_ch + self.pool = nn.AdaptiveAvgPool2d((1, 1)) + self.net = nn.Sequential( + nn.Linear(in_ch, hidden), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + nn.Linear(hidden, num_label_classes), + ) + + def forward(self, feat: torch.Tensor, logits_seg: torch.Tensor) -> torch.Tensor: + feat = F.interpolate( + feat, + size=logits_seg.shape[-2:], + mode="bilinear", + align_corners=False, + ) + x = torch.cat([feat, logits_seg], dim=1) + x = self.pool(x).flatten(1) + return self.net(x) + + +def get_last_feat(out, logits: torch.Tensor) -> torch.Tensor: + if hasattr(out, "hidden_states") and out.hidden_states is not None: + feat = out.hidden_states[-1] + else: + feat = logits + + feat = F.interpolate( + feat, + size=logits.shape[-2:], + mode="bilinear", + align_corners=False, + ) + return feat + + +class VisualSegformerDualLabel(nn.Module): + def __init__(self, base_model: nn.Module, label_head: nn.Module): + super().__init__() + self.base_model = base_model + self.label_head = label_head + + def forward(self, pixel_values: torch.Tensor) -> Dict[str, torch.Tensor]: + out = self.base_model(pixel_values=pixel_values) + logits_seg = out.logits + feat = get_last_feat(out, logits_seg) + logits_label = self.label_head(feat, logits_seg) + return { + "semantic_logits": logits_seg, + "label_logits": logits_label, + "label_probs": torch.softmax(logits_label, dim=1), + } + + +def build_visual_model( + backbone: str, + num_seg_classes: int, + num_label_classes: int, + device: torch.device, + input_hw: Tuple[int, int], +) -> VisualSegformerDualLabel: + H, W = input_hw + + base_model = SegformerForSemanticSegmentation.from_pretrained( + backbone, + num_labels=int(num_seg_classes), + ignore_mismatched_sizes=True, + use_safetensors=True, + ) + base_model.config.output_hidden_states = True + base_model.to(device) + base_model.eval() + + with torch.no_grad(): + dummy = torch.zeros((1, 3, int(H), int(W)), dtype=torch.float32, device=device) + out = base_model(pixel_values=dummy) + logits = out.logits + feat = get_last_feat(out, logits) + feat_ch = int(feat.shape[1]) + + label_head = LabelHead( + feat_ch=feat_ch, + num_seg_classes=int(num_seg_classes), + num_label_classes=int(num_label_classes), + hidden=256, + dropout=0.2, + ).to(device) + label_head.eval() + + return VisualSegformerDualLabel(base_model=base_model, label_head=label_head).to(device) + + +def load_checkpoint_into_model(model: VisualSegformerDualLabel, checkpoint_path: Path): + ckpt = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + + if "model" not in ckpt: + raise RuntimeError(f"Checkpoint não contém chave 'model': {checkpoint_path}") + + if "aux_head" not in ckpt: + raise RuntimeError(f"Checkpoint não contém chave 'aux_head': {checkpoint_path}") + + model.base_model.load_state_dict(ckpt["model"], strict=True) + model.label_head.load_state_dict(ckpt["aux_head"], strict=True) + return ckpt + + +class VisualTorchTupleWrapper(nn.Module): + def __init__(self, model: VisualSegformerDualLabel, output_kind: str = "contract"): + super().__init__() + self.model = model + self.output_kind = str(output_kind).lower() + + if self.output_kind not in ("raw", "contract"): + raise RuntimeError(f"output_kind inválido: {self.output_kind}") + + def forward(self, pixel_values: torch.Tensor): + outputs = self.model(pixel_values=pixel_values) + semantic = outputs["semantic_logits"] + label_probs = outputs["label_probs"] + + if self.output_kind == "contract": + semantic = F.interpolate( + semantic, + size=pixel_values.shape[-2:], + mode="bilinear", + align_corners=False, + ) + + return semantic, label_probs + + +# ============================================================ +# Benchmark PyTorch +# ============================================================ + + +@torch.inference_mode() +def benchmark_torch( + model: nn.Module, + inputs_np: List[np.ndarray], + device: torch.device, + warmup: int, + repeat: int, + amp: bool, + label: str, +) -> Tuple[dict, List[dict]]: + model.eval() + + times: List[float] = [] + rows: List[dict] = [] + + inputs_t = [ + torch.from_numpy(x).to(device, non_blocking=True) + for x in inputs_np + ] + + synchronize_if_cuda(device) + + print(f"\n[BENCH] {label} | warmup={warmup} repeat={repeat}") + + for i in range(max(0, warmup)): + x = inputs_t[i % len(inputs_t)] + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=amp and device.type == "cuda"): + _ = model(x) + + synchronize_if_cuda(device) + + total_iter = len(inputs_t) * max(1, repeat) + idx = 0 + + for r in range(max(1, repeat)): + for sample_idx, x in enumerate(inputs_t): + synchronize_if_cuda(device) + t0 = time.perf_counter() + + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=amp and device.type == "cuda"): + _ = model(x) + + synchronize_if_cuda(device) + dt_ms = (time.perf_counter() - t0) * 1000.0 + + times.append(dt_ms) + rows.append({ + "engine": label, + "repeat": r, + "sample_idx": sample_idx, + "iter_idx": idx, + "latency_ms": dt_ms, + }) + + idx += 1 + if idx % 25 == 0 or idx == total_iter: + print(f" {idx:04d}/{total_iter:04d} | last={dt_ms:.2f}ms") + + summary = summarize_times(times) + summary["engine"] = label + return summary, rows + + +# ============================================================ +# ONNX Runtime +# ============================================================ + + +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 CPU:\n" + " pip install onnxruntime" + ) + + provider = provider.lower() + + 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")) + + dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.4\bin") + dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.3\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), + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str(cache_dir), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(cache_dir), + "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}" + ) + + session = ort.InferenceSession( + str(onnx_path), + sess_options=sess_options, + providers=providers_ok, + ) + + active = session.get_providers() + print(f"[ONNX] usando providers: {active}") + + if provider == "tensorrt" and "TensorrtExecutionProvider" not in active: + raise RuntimeError( + "TensorRTExecutionProvider foi solicitado, mas não ficou ativo. " + f"Providers ativos: {active}." + ) + + if provider == "cuda" and "CUDAExecutionProvider" not in active: + raise RuntimeError( + "CUDAExecutionProvider foi solicitado, mas não ficou ativo. " + f"Providers ativos: {active}." + ) + + return session + + +def benchmark_onnx( + session, + inputs_np: List[np.ndarray], + warmup: int, + repeat: int, + label: str, +) -> Tuple[dict, List[dict]]: + input_name = session.get_inputs()[0].name + + times: List[float] = [] + rows: List[dict] = [] + + print(f"\n[BENCH] {label} | warmup={warmup} repeat={repeat}") + + for i in range(max(0, warmup)): + x = inputs_np[i % len(inputs_np)] + _ = session.run(None, {input_name: x}) + + total_iter = len(inputs_np) * max(1, repeat) + idx = 0 + + for r in range(max(1, repeat)): + for sample_idx, x in enumerate(inputs_np): + t0 = time.perf_counter() + _ = session.run(None, {input_name: x}) + dt_ms = (time.perf_counter() - t0) * 1000.0 + + times.append(dt_ms) + rows.append({ + "engine": label, + "repeat": r, + "sample_idx": sample_idx, + "iter_idx": idx, + "latency_ms": dt_ms, + }) + + idx += 1 + if idx % 25 == 0 or idx == total_iter: + print(f" {idx:04d}/{total_iter:04d} | last={dt_ms:.2f}ms") + + summary = summarize_times(times) + summary["engine"] = label + return summary, rows + + +# ============================================================ +# 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("--labelmap", default="") + parser.add_argument("--norm_stats", default="") + + parser.add_argument("--split_folder", default="val", choices=["train", "val", "test"]) + parser.add_argument("--root_override", default=None) + parser.add_argument("--test_folder", default=None) + + parser.add_argument("--max_samples", type=int, default=50) + parser.add_argument("--start_idx", type=int, default=0) + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--repeat", type=int, default=5) + + 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") + parser.add_argument("--trt_no_fp16", action="store_true") + + parser.add_argument("--skip_torch_fp32", action="store_true") + parser.add_argument("--skip_torch_amp", action="store_true") + parser.add_argument("--skip_onnx", action="store_true") + + parser.add_argument( + "--onnx_has_norm", + action="store_true", + help="Use quando o ONNX já inclui normalização interna. Nesse caso o ONNX recebe RGB 0..1 cru.", + ) + + parser.add_argument( + "--torch_output_kind", + default="contract", + choices=["raw", "contract"], + help="contract mede PyTorch com semantic_logits redimensionado para HxW, igual ao ONNX exportado com resize_logits.", + ) + + parser.add_argument("--out_dir", default=None) + + 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) + + checkpoint_path, onnx_path, ckpt_name, mode, save_dir = resolve_default_paths(args, config, config_dir) + + if mode != "label": + raise RuntimeError(f"Este benchmark foi preparado para dual_head_label. Modo detectado: {mode}") + + labelmap_path = resolve_labelmap_path(args, config, config_dir) + semantic_id2label, semantic_label2id, ignore_index = load_labelmap(labelmap_path) + num_seg_classes = len(semantic_id2label) + + W, H = config.get("resolucao", [1024, 640]) + W = int(W) + H = int(H) + backbone = str(config.get("backbone", "nvidia/mit-b0")) + + ckpt_meta = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + label_name_by_id, num_label_classes = resolve_label_classes(config, ckpt_meta) + + norm_stats_path = resolve_norm_stats_path(args, config, config_dir, save_dir) + mean, std, norm_stats_used, norm_channels = load_rgb_norm_stats(norm_stats_path) + + if args.test_folder: + root = resolve_path(args.test_folder, Path.cwd()) + if root is None or not root.is_dir(): + raise FileNotFoundError(f"Pasta de teste não encontrada: {root}") + samples = discover_image_folder( + folder=root, + max_samples=args.max_samples, + start_idx=args.start_idx, + ) + else: + if args.root_override: + root = resolve_path(args.root_override, Path.cwd()) + else: + camera = str(config.get("camera", "oak-d")) + root = (config_dir / camera / "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 = discover_samples( + split_root=root, + max_samples=args.max_samples, + start_idx=args.start_idx, + ) + + if args.out_dir: + out_dir = resolve_path(args.out_dir, Path.cwd()) + else: + out_dir = onnx_path.parent / "benchmarks" + + assert out_dir is not None + out_dir.mkdir(parents=True, exist_ok=True) + + 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("Benchmark Visual Worker PyTorch vs ONNX") + print(f"Config : {config_path}") + print(f"Checkpoint : {checkpoint_path}") + print(f"ONNX : {onnx_path}") + print(f"Root : {root}") + print(f"Samples : {len(samples)}") + print(f"Warmup : {args.warmup}") + print(f"Repeat : {args.repeat}") + print(f"Backbone : {backbone}") + print(f"Input shape : [1, 3, {H}, {W}]") + print(f"Semantic classes: {num_seg_classes} {semantic_id2label}") + print(f"Label classes : {num_label_classes} {label_name_by_id}") + print(f"Device : {device}") + print(f"ONNX provider : {args.onnx_provider}") + print(f"ONNX has norm : {args.onnx_has_norm}") + print(f"Torch output : {args.torch_output_kind}") + print(f"Out dir : {out_dir}") + print("==========================================") + + # Entradas PyTorch: sempre normalizadas, porque o modelo PyTorch puro espera normalização externa. + # Entradas ONNX: se ONNX tem norm embutida, entram 0..1; senão, entram normalizadas. + print("\n[DATA] Carregando imagens na RAM...") + inputs_torch_np = load_inputs_as_numpy( + samples=samples, + mean=mean, + std=std, + target_hw=(H, W), + normalize_input=True, + ) + + if args.onnx_has_norm: + inputs_onnx_np = load_inputs_as_numpy( + samples=samples, + mean=mean, + std=std, + target_hw=(H, W), + normalize_input=False, + ) + print("[DATA] ONNX receberá RGB 0..1 cru, pois --onnx_has_norm está ativo.") + else: + inputs_onnx_np = inputs_torch_np + print("[DATA] ONNX receberá RGB normalizado, pois --onnx_has_norm não está ativo.") + + print(f"[DATA] Inputs carregados: {len(inputs_torch_np)}") + + summaries: List[dict] = [] + all_rows: List[dict] = [] + + # ======================================================== + # PyTorch + # ======================================================== + need_torch = not args.skip_torch_fp32 or not args.skip_torch_amp + + if need_torch: + print("\n[MODEL] Montando PyTorch...") + model = build_visual_model( + backbone=backbone, + num_seg_classes=num_seg_classes, + num_label_classes=num_label_classes, + device=device, + input_hw=(H, W), + ) + + print("[CKPT] Carregando checkpoint...") + ckpt = load_checkpoint_into_model(model, checkpoint_path) + model.to(device) + model.eval() + + torch_model = VisualTorchTupleWrapper( + model=model, + output_kind=args.torch_output_kind, + ).to(device) + torch_model.eval() + + clear_cuda() + + if not args.skip_torch_fp32: + summary, rows = benchmark_torch( + model=torch_model, + inputs_np=inputs_torch_np, + device=device, + warmup=args.warmup, + repeat=args.repeat, + amp=False, + label="torch_fp32", + ) + summaries.append(summary) + all_rows.extend(rows) + + clear_cuda() + + if not args.skip_torch_amp: + summary, rows = benchmark_torch( + model=torch_model, + inputs_np=inputs_torch_np, + device=device, + warmup=args.warmup, + repeat=args.repeat, + amp=True, + label="torch_amp_fp16", + ) + summaries.append(summary) + all_rows.extend(rows) + + del torch_model + del model + clear_cuda() + + # ======================================================== + # ONNX + # ======================================================== + if not args.skip_onnx: + print("\n[ONNX] Carregando sessão...") + session = create_onnx_session( + onnx_path=onnx_path, + provider=args.onnx_provider, + trt_home=args.trt_home, + trt_fp16=not args.trt_no_fp16, + ) + + summary, rows = benchmark_onnx( + session=session, + inputs_np=inputs_onnx_np, + warmup=args.warmup, + repeat=args.repeat, + label=f"onnx_{args.onnx_provider}", + ) + summaries.append(summary) + all_rows.extend(rows) + + # ======================================================== + # Relatório + # ======================================================== + print("\n========== RESUMO ==========") + + for s in summaries: + print( + f"{s['engine']:<16} " + f"n={s['n']:<5} " + f"mean={s['mean_ms']:.3f}ms " + f"median={s['median_ms']:.3f}ms " + f"p95={s['p95_ms']:.3f}ms " + f"p99={s['p99_ms']:.3f}ms " + f"fps_mean={s['fps_mean']:.2f} " + f"fps_p95={s['fps_p95_latency']:.2f}" + ) + + base_name = f"{onnx_path.stem}_{args.onnx_provider}" + report_json = out_dir / f"{base_name}_visual_benchmark_report.json" + report_csv = out_dir / f"{base_name}_visual_benchmark_rows.csv" + + report = { + "kind": "visual_worker_benchmark_onnx", + "config": str(config_path), + "checkpoint": str(checkpoint_path), + "ckpt_name": ckpt_name, + "onnx": str(onnx_path), + "onnx_has_norm": bool(args.onnx_has_norm), + "root": str(root), + "samples": len(samples), + "warmup": int(args.warmup), + "repeat": int(args.repeat), + "input_shape": [1, 3, H, W], + "input_channel_names": ["R", "G", "B"], + "semantic_id2label": semantic_id2label, + "label_name_by_id": label_name_by_id, + "norm_stats_used": norm_stats_used, + "norm_channels": norm_channels, + "onnx_provider": args.onnx_provider, + "trt_fp16": bool(not args.trt_no_fp16), + "device": str(device), + "torch_output_kind": args.torch_output_kind, + "summaries": summaries, + "samples_list": [ + { + "img_path": s.img_path, + "group_name": s.group_name, + "filename": s.filename, + } + for s in samples + ], + } + + save_json(report_json, report) + save_csv(report_csv, all_rows) + + print(f"\n[OK] JSON salvo em: {report_json}") + print(f"[OK] CSV salvo em : {report_csv}") + print("\nBenchmark finalizado.") + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/_9_test_segformer_dual.py b/Python/OAK/datasets/_9_test_segformer_dual.py index 9cfa42e0a..2e14860a4 100644 --- a/Python/OAK/datasets/_9_test_segformer_dual.py +++ b/Python/OAK/datasets/_9_test_segformer_dual.py @@ -49,10 +49,14 @@ from transformers import SegformerForSemanticSegmentation # Utils básicos # ============================================================ -IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1) -IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1) +# Normalização padrão. Pode ser sobrescrita por norm_stats.json. +NORM_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1) +NORM_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1) + def normalize_img(img: torch.Tensor) -> torch.Tensor: - return (img - IMAGENET_MEAN.to(img.device)) / IMAGENET_STD.to(img.device) + mean = NORM_MEAN.to(img.device) + std = NORM_STD.to(img.device).clamp_min(1e-6) + return (img - mean) / std def find_device(): @@ -477,6 +481,229 @@ def infer_sample(base_model, aux_head, img_bgr: np.ndarray, device: torch.device return result +def add_runtime_dll_dirs(trt_home: Optional[str] = None): + 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")) + + dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.4\bin") + dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.3\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}") + + +def create_onnx_session(onnx_path: str, provider: str = "tensorrt", trt_home: Optional[str] = None): + try: + import onnxruntime as ort + except ImportError: + raise ImportError( + "onnxruntime não está instalado. Instale com:\n" + " pip install onnxruntime-gpu\n" + "ou CPU:\n" + " pip install onnxruntime" + ) + + provider = provider.lower() + + if provider == "tensorrt": + add_runtime_dll_dirs(trt_home) + + 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 = os.path.join(os.path.dirname(onnx_path), "trt_cache") + os.makedirs(cache_dir, exist_ok=True) + + trt_options = { + "device_id": 0, + "trt_fp16_enable": True, + "trt_engine_cache_enable": True, + "trt_engine_cache_path": cache_dir, + "trt_timing_cache_enable": True, + "trt_timing_cache_path": cache_dir, + "trt_max_workspace_size": 4 * 1024 * 1024 * 1024, + } + + providers = [ + ("TensorrtExecutionProvider", trt_options), + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + else: + raise RuntimeError(f"Provider desconhecido: {provider}") + + 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. disponível={available}") + + sess = ort.InferenceSession( + onnx_path, + sess_options=sess_options, + providers=providers_ok, + ) + + active = sess.get_providers() + print(f"[ONNX] usando providers: {active}") + + if provider == "tensorrt" and "TensorrtExecutionProvider" not in active: + raise RuntimeError(f"TensorRT solicitado, mas não ficou ativo. Providers ativos: {active}") + + if provider == "cuda" and "CUDAExecutionProvider" not in active: + raise RuntimeError(f"CUDA solicitado, mas não ficou ativo. Providers ativos: {active}") + + return sess + + +def preprocess_img_onnx(img_bgr: np.ndarray, resolucao, onnx_has_norm: bool): + """ + Retorna input NCHW float32. + + Se ONNX tem normalização embutida: + envia RGB 0..1 + + Se ONNX NÃO tem normalização embutida: + envia RGB já normalizado com NORM_MEAN/NORM_STD + """ + w, h = int(resolucao[0]), int(resolucao[1]) + + img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB) + img_res = cv2.resize(img_rgb, (w, h), interpolation=cv2.INTER_AREA) + + x = img_res.astype(np.float32) / 255.0 + x = np.transpose(x, (2, 0, 1)) # CHW + + if not onnx_has_norm: + mean = NORM_MEAN.detach().cpu().numpy().astype(np.float32) + std = NORM_STD.detach().cpu().numpy().astype(np.float32) + std = np.clip(std, 1e-6, None) + x = (x - mean) / std + + return np.expand_dims(x, axis=0).astype(np.float32) + + +def get_onnx_outputs(outputs_dict: Dict[str, np.ndarray]): + keys = list(outputs_dict.keys()) + + semantic_name = None + label_name = None + + for k in ["semantic_logits", "semantic_mask", "semantic", "output_semantic"]: + if k in outputs_dict: + semantic_name = k + break + + for k in ["label_probs", "label_logits", "label", "output_label"]: + if k in outputs_dict: + label_name = k + break + + if semantic_name is None or label_name is None: + if len(keys) != 2: + raise RuntimeError(f"Outputs ONNX inesperados: {keys}") + semantic_name = semantic_name or keys[0] + label_name = label_name or keys[1] + + return outputs_dict[semantic_name], outputs_dict[label_name], semantic_name, label_name + + +def infer_sample_onnx( + onnx_session, + img_bgr: np.ndarray, + mode: str, + resolucao, + onnx_has_norm: bool, + mask2_thr: float = 0.5, +): + h0, w0 = img_bgr.shape[:2] + + x = preprocess_img_onnx( + img_bgr=img_bgr, + resolucao=resolucao, + onnx_has_norm=onnx_has_norm, + ) + + input_name = onnx_session.get_inputs()[0].name + output_names = [o.name for o in onnx_session.get_outputs()] + raw_outputs = onnx_session.run(None, {input_name: x}) + + outputs_dict = { + name: arr.astype(np.float32) + for name, arr in zip(output_names, raw_outputs) + } + + semantic_out, label_out, semantic_name, label_name = get_onnx_outputs(outputs_dict) + + result = { + "pred_seg": None, + "mask2_prob": None, + "mask2_bin": None, + "label_id": None, + "label_conf": None, + "label_probs": None, + } + + # semantic_logits: [1,C,H,W] + # semantic_mask : [1,H,W] + if semantic_out.ndim == 4: + pred_ids = np.argmax(semantic_out, axis=1)[0].astype(np.uint8) + elif semantic_out.ndim == 3: + pred_ids = semantic_out[0].astype(np.uint8) + else: + raise RuntimeError(f"Saída semântica ONNX inválida: {semantic_name} shape={semantic_out.shape}") + + pred_ids_full = cv2.resize(pred_ids, (w0, h0), interpolation=cv2.INTER_NEAREST) + result["pred_seg"] = pred_ids_full + + # Neste modelo atual não estamos usando mask2 no ONNX. + if mode == "mask2": + result["mask2_prob"] = None + result["mask2_bin"] = None + return result + + if mode == "label": + if label_name == "label_probs": + probs = label_out[0].astype(np.float32) + else: + logits = label_out.astype(np.float32) + logits = logits - np.max(logits, axis=1, keepdims=True) + e = np.exp(logits) + probs = (e / np.clip(np.sum(e, axis=1, keepdims=True), 1e-12, None))[0] + + lid = int(np.argmax(probs)) + result["label_id"] = lid + result["label_conf"] = float(probs[lid]) + result["label_probs"] = probs + + return result + + # ============================================================ # Painéis # ============================================================ @@ -636,6 +863,10 @@ def main(): ap.add_argument("--camera", action="store_true", help="Inferir em tempo real pela OAK-D Lite") ap.add_argument("--camera_fps", type=int, default=20) ap.add_argument("--camera_res", type=str, default="720p", choices=["720p", "1080p"]) + ap.add_argument("--onnx", type=str, default=None, help="Se informado, usa ONNX Runtime em vez de PyTorch.") + ap.add_argument("--onnx_provider", type=str, default="tensorrt", choices=["cuda", "cpu", "tensorrt"]) + ap.add_argument("--onnx_has_norm", action="store_true", help="Use se o ONNX foi exportado com --include-norm.") + ap.add_argument("--trt_home", type=str, default=r"C:\dev\TensorRT-10.10.0.31") args = ap.parse_args() device = find_device() @@ -690,26 +921,36 @@ def main(): raise SystemExit(f"Checkpoint não encontrado: {ckpt_path}") print(f"[model] ckpt={ckpt_path}") - base_model = build_base_model(num_classes=num_classes, device=device, backbone=BACKBONE) + use_onnx = args.onnx is not None and str(args.onnx).strip() != "" + onnx_session = None + base_model = None aux_head = None label_names = {} - feat_ch = get_feat_ch(base_model, device, input_h=args.input_size, input_w=args.input_size) - # Carrega ckpt antes para descobrir metadados se necessário - ckpt_pre = torch.load(ckpt_path, map_location="cpu", weights_only=False) - extra = ckpt_pre.get("extra", {}) or {} + if use_onnx: + if not os.path.isfile(args.onnx): + raise SystemExit(f"ONNX não encontrado: {args.onnx}") - if mode == "mask2": - aux_head = Mask2Head(feat_ch=feat_ch, num_classes=num_classes, hidden=256, dropout=0.1) - elif mode == "label": - print(extra) - label_names_raw = extra.get("label_name_by_id", {}) or {} - # JSON pode salvar keys como string - label_names = {int(k): str(v) for k, v in label_names_raw.items()} if label_names_raw else {} - print(label_names_raw) + print(f"[runtime] usando ONNX: {args.onnx}") + print(f"[runtime] provider: {args.onnx_provider}") + print(f"[runtime] onnx_has_norm: {args.onnx_has_norm}") - # Se não tiver no ckpt, tenta inferir do dataset de labels + onnx_session = create_onnx_session( + onnx_path=args.onnx, + provider=args.onnx_provider, + trt_home=args.trt_home, + ) + + # Ainda precisamos dos nomes dos labels para HUD. + # Primeiro tenta buscar no checkpoint, se existir. + if os.path.isfile(ckpt_path): + ckpt_pre = torch.load(ckpt_path, map_location="cpu", weights_only=False) + extra = ckpt_pre.get("extra", {}) or {} + label_names_raw = extra.get("label_name_by_id", {}) or {} + label_names = {int(k): str(v) for k, v in label_names_raw.items()} if label_names_raw else {} + + # Fallback: tenta inferir dos labels do dataset. max_label_id = -1 for s in samples: lid, lname = read_gt_label(s) @@ -717,23 +958,94 @@ def main(): max_label_id = max(max_label_id, int(lid)) if lname: label_names[int(lid)] = str(lname) - if max_label_id < 0: - # último fallback: pelo shape da última camada do checkpoint + + if mode == "label": + if max_label_id < 0 and label_names: + max_label_id = max(label_names.keys()) + + for i in range(max_label_id + 1): + label_names.setdefault(i, f"label_{i}") + + # Fallback final para seu caso conhecido. + if not label_names: + label_names = { + 0: "Parado", + 1: "EntrandoRua", + 2: "CaminhandoRua", + 3: "SaindoRua", + 4: "Manobrando", + 5: "Direcionando", + 6: "RetornandoBase", + 7: "Indefinido", + } + + print(f"[label_names] {label_names}") + + else: + print(f"[runtime] usando PyTorch: {ckpt_path}") + + base_model = build_base_model(num_classes=num_classes, device=device, backbone=BACKBONE) + + aux_head = None + label_names = {} + feat_ch = get_feat_ch(base_model, device, input_h=args.input_size, input_w=args.input_size) + + ckpt_pre = torch.load(ckpt_path, map_location="cpu", weights_only=False) + extra = ckpt_pre.get("extra", {}) or {} + + if mode == "mask2": + aux_head = Mask2Head(feat_ch=feat_ch, num_classes=num_classes, hidden=256, dropout=0.1) + + elif mode == "label": + label_names_raw = extra.get("label_name_by_id", {}) or {} + label_names = {int(k): str(v) for k, v in label_names_raw.items()} if label_names_raw else {} + + max_label_id = -1 + for s in samples: + lid, lname = read_gt_label(s) + if lid is not None: + max_label_id = max(max_label_id, int(lid)) + if lname: + label_names[int(lid)] = str(lname) + + # Fonte mais confiável: shape da última camada salva no checkpoint. + ckpt_num_label_classes = None sd = ckpt_pre.get("aux_head", {}) + for k, v in sd.items(): if k.endswith("net.3.weight") or k.endswith("net.3.bias"): - max_label_id = int(v.shape[0]) - 1 + ckpt_num_label_classes = int(v.shape[0]) break - num_label_classes = max_label_id + 1 - if num_label_classes <= 0: - raise RuntimeError("Não consegui inferir num_label_classes para LabelHead.") - for i in range(num_label_classes): - label_names.setdefault(i, f"label_{i}") - aux_head = LabelHead(feat_ch=feat_ch, num_seg_classes=num_classes, num_label_classes=num_label_classes, hidden=256, dropout=0.2) - print(f"[label_head] classes={num_label_classes} names={label_names}") - ckpt = load_checkpoint(ckpt_path, base_model, aux_head, device) - print(f"[ckpt] epoch={ckpt.get('epoch')} bests={ckpt.get('bests')}") + if ckpt_num_label_classes is not None and ckpt_num_label_classes > 0: + num_label_classes = ckpt_num_label_classes + else: + # Fallback: labels vistos no dataset atual. + num_label_classes = max_label_id + 1 + + if num_label_classes <= 0: + raise RuntimeError("Não consegui inferir num_label_classes para LabelHead.") + + for i in range(num_label_classes): + label_names.setdefault(i, f"label_{i}") + if num_label_classes <= 0: + raise RuntimeError("Não consegui inferir num_label_classes para LabelHead.") + + for i in range(num_label_classes): + label_names.setdefault(i, f"label_{i}") + + aux_head = LabelHead( + feat_ch=feat_ch, + num_seg_classes=num_classes, + num_label_classes=num_label_classes, + hidden=256, + dropout=0.2, + ) + + print(f"[label_head] classes={num_label_classes} names={label_names}") + + ckpt = load_checkpoint(ckpt_path, base_model, aux_head, device) + print(f"[ckpt] epoch={ckpt.get('epoch')} bests={ckpt.get('bests')}") if args.save_dir: os.makedirs(args.save_dir, exist_ok=True) @@ -809,15 +1121,25 @@ def main(): roi_bgr = frame_bgr[y0:y1, 0:W] roi_input = resize_to_config(roi_bgr, RESOLUCAO) - res = infer_sample( - base_model=base_model, - aux_head=aux_head, - img_bgr=roi_input, - device=device, - mode=mode, - input_size=args.input_size, - mask2_thr=args.mask2_thr, - ) + if use_onnx: + res = infer_sample_onnx( + onnx_session=onnx_session, + img_bgr=roi_input, + mode=mode, + resolucao=RESOLUCAO, + onnx_has_norm=args.onnx_has_norm, + mask2_thr=args.mask2_thr, + ) + else: + res = infer_sample( + base_model=base_model, + aux_head=aux_head, + img_bgr=roi_input, + device=device, + mode=mode, + input_size=args.input_size, + mask2_thr=args.mask2_thr, + ) pred_seg = res["pred_seg"] pred_col = colorize_seg(pred_seg, lut_bgr) @@ -896,15 +1218,25 @@ def main(): gt2 = imread_gray(s.mask2_path) if s.mask2_path else None gt_label_id, gt_label_name = read_gt_label(s) - res = infer_sample( - base_model=base_model, - aux_head=aux_head, - img_bgr=img, - device=device, - mode=mode, - input_size=args.input_size, - mask2_thr=args.mask2_thr, - ) + if use_onnx: + res = infer_sample_onnx( + onnx_session=onnx_session, + img_bgr=img, + mode=mode, + resolucao=RESOLUCAO, + onnx_has_norm=args.onnx_has_norm, + mask2_thr=args.mask2_thr, + ) + else: + res = infer_sample( + base_model=base_model, + aux_head=aux_head, + img_bgr=img, + device=device, + mode=mode, + input_size=args.input_size, + mask2_thr=args.mask2_thr, + ) pred1 = res["pred_seg"] segm = seg_metrics(pred1, gt1, num_classes=num_classes, ignore_id=ignore_id) diff --git a/Python/OAK/datasets/config.json b/Python/OAK/datasets/config.json index c87c21af3..0097dfa44 100644 --- a/Python/OAK/datasets/config.json +++ b/Python/OAK/datasets/config.json @@ -7,6 +7,7 @@ "main_class_name": "navegavel", "es_classes": "", "model_to_use": "geral", + "ckpt_test": "best_main", "raw_size": [1296, 1028], "resolucao": [1024, 576], "roi_inicio": 0.0, diff --git a/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py b/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py new file mode 100644 index 000000000..693c1912f --- /dev/null +++ b/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py @@ -0,0 +1,508 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +_10_export_onnx.py + +Exporta o checkpoint PyTorch do SegFormer Multi-Head OAK-FCC-3 para ONNX. + +Exemplo: + +python _10_export_onnx.py --config config.json --checkpoint backup/segformer_b1/target_teached/stacked_raw5/best_score.pt --out backup/segformer_b1/target_teached/stacked_raw5/best_score.onnx --train-script _8_train_multihead.py --opset 17 --device cuda +""" + +from __future__ import annotations + +import os +import json +import argparse +import importlib.util +from pathlib import Path +from typing import Dict, List + +import torch +import torch.nn as nn + + +# ============================================================ +# Utils +# ============================================================ + +def load_json(path: str | Path) -> dict: + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + + +def ensure_dir(path: Path): + path.mkdir(parents=True, exist_ok=True) + + +def import_train_module(train_script_path: str | Path): + """ + Importa o script de treinamento como módulo, para reaproveitar: + - load_labelmap + - build_heads_config + - get_input_channel_names + - get_input_channel_indices + - build_model + - load_checkpoint + """ + 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 save_export_metadata(path: Path, data: dict): + ensure_dir(path.parent) + with path.open("w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + + +def resolve_norm_stats_path(config: dict, config_path: Path) -> Path | None: + fusion_mode = config.get("fusion_mode", "stacked") + model = config.get("modelo") + model_name = config.get("model_name") + ch = config.get("channels") + + normstats_path = Path(f"backup/{model}/{model_name}/{fusion_mode}_raw{ch}/norm_stats.json") + if normstats_path.exists(): + return normstats_path + + ia_resolution = config.get("resolucao", [1024, 640]) + w, h = int(ia_resolution[0]), int(ia_resolution[1]) + p = Path("dataset") / f"{w}x{h}" / "group" / "norm_stats.json" + if not p.is_absolute(): + p = (config_path.parent / p).resolve() + return p + + + return None + + +def load_selected_norm_stats(config: dict, config_path: Path, input_channel_indices: List[int]): + path = resolve_norm_stats_path(config, config_path) + + if path is None or not path.is_file(): + raise FileNotFoundError(f"norm_stats não encontrado: {path}") + + js = load_json(path) + + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + max_idx = max(input_channel_indices) + + if len(mean) <= max_idx or len(std) <= max_idx: + raise RuntimeError( + f"norm_stats incompatível: precisa índices={input_channel_indices}, " + f"mean={len(mean)} std={len(std)} path={path}" + ) + + mean_sel = [float(mean[i]) for i in input_channel_indices] + std_sel = [float(std[i]) for i in input_channel_indices] + + if names: + names_sel = [names[i] for i in input_channel_indices] + else: + names_sel = [] + + return path, mean_sel, std_sel, names_sel + + +# ============================================================ +# Wrapper ONNX +# ============================================================ + +class MultiHeadOnnxWrapper(nn.Module): + """ + Wrapper ONNX para exportar o SegFormer Multi-Head. + + Pode exportar: + - logits crus + - logits redimensionados + - argmax em baixa resolução + - argmax em resolução da entrada + + Também pode embutir a normalização: + x = (x - mean) / std + """ + + def __init__( + self, + model: nn.Module, + output_heads: List[str], + resize_to_input: bool = False, + include_norm: bool = False, + norm_mean: List[float] | None = None, + norm_std: List[float] | None = None, + postprocess: str = "none", + ): + super().__init__() + self.model = model + self.output_heads = list(output_heads) + self.resize_to_input = bool(resize_to_input) + self.include_norm = bool(include_norm) + 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}") + + if self.postprocess == "resize_logits": + self.resize_to_input = True + + if self.include_norm: + if norm_mean is None or norm_std is None: + raise RuntimeError("include_norm=True requer norm_mean e norm_std") + + mean_t = torch.tensor(norm_mean, dtype=torch.float32).view(1, len(norm_mean), 1, 1) + std_t = torch.tensor(norm_std, dtype=torch.float32).view(1, len(norm_std), 1, 1) + + self.register_buffer("norm_mean", mean_t) + self.register_buffer("norm_std", std_t) + else: + self.register_buffer("norm_mean", torch.empty(0)) + self.register_buffer("norm_std", torch.empty(0)) + + def forward(self, pixel_values: torch.Tensor): + input_hw = pixel_values.shape[-2:] + + x = pixel_values + + if self.include_norm: + x = (x - self.norm_mean) / torch.clamp(self.norm_std, min=1e-6) + + outputs: Dict[str, torch.Tensor] = self.model(pixel_values=x) + + result = [] + + for head_name in self.output_heads: + if head_name not in outputs: + raise RuntimeError(f"Head ausente no modelo: {head_name}") + + logits = outputs[head_name] + + if self.postprocess in ("none", "resize_logits"): + if self.resize_to_input: + logits = torch.nn.functional.interpolate( + logits, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + + result.append(logits) + + elif self.postprocess == "argmax_lowres": + mask = torch.argmax(logits, dim=1).to(torch.uint8) + result.append(mask) + + elif self.postprocess == "argmax_fullres": + logits = torch.nn.functional.interpolate( + logits, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + mask = torch.argmax(logits, dim=1).to(torch.uint8) + result.append(mask) + + return tuple(result) + + +# ============================================================ +# Main +# ============================================================ + +def main(): + parser = argparse.ArgumentParser() + + parser.add_argument("--config", default="config.json") + parser.add_argument("--checkpoint", default="") + parser.add_argument("--out", default="") + + parser.add_argument( + "--train-script", + default="_8_train_multihead.py", + help="Script de treino usado para montar exatamente o mesmo modelo.", + ) + + parser.add_argument("--labelmap", default="dataset/labelmap.txt") + parser.add_argument("--opset", type=int, default=17) + parser.add_argument("--device", default="cuda", choices=["cuda", "cpu"]) + + parser.add_argument( + "--resize-to-input", + action="store_true", + help=( + "Se ativo, exporta cada head já redimensionada para HxW da entrada. " + "Se desativo, exporta logits brutos do decode_head." + ), + ) + + parser.add_argument( + "--include-norm", + action="store_true", + help="Inclui normalização x=(x-mean)/std dentro do grafo ONNX.", + ) + + parser.add_argument( + "--postprocess", + default="none", + choices=["none", "resize_logits", "argmax_lowres", "argmax_fullres"], + help=( + "Define o pós-processamento exportado no ONNX. " + "none=logits crus; resize_logits=logits em HxW; " + "argmax_lowres=máscara na resolução do decode_head; " + "argmax_fullres=máscara HxW pronta." + ), + ) + + parser.add_argument( + "--dynamic-batch", + action="store_true", + help="Permite batch dinâmico no ONNX. H e W continuam fixos.", + ) + + args = parser.parse_args() + + config_path = Path(args.config) + checkpoint_path = Path(args.checkpoint) + out_path = Path(args.out) + labelmap_path = Path(args.labelmap) + + if not config_path.exists(): + raise FileNotFoundError(f"Config não encontrado: {config_path}") + + if args.checkpoint != "" and not checkpoint_path.exists(): + raise FileNotFoundError(f"Checkpoint não encontrado: {checkpoint_path}") + + if not labelmap_path.exists(): + raise FileNotFoundError(f"Labelmap não encontrado: {labelmap_path}") + + if args.out != "": + ensure_dir(out_path.parent) + + train_mod = import_train_module(args.train_script) + + config = load_json(config_path) + + W, H = config["resolucao"] + backbone = config.get("backbone", "nvidia/mit-b1") + fusion_mode = config.get("fusion_mode", "stacked") + model = config.get("modelo") + model_name = config.get("model_name") + ch = config.get("channels") + ckpt_name = config.get("ckpt_test", "best_score") + + if args.checkpoint == "": + checkpoint_path = Path(f"backup/{model}/{model_name}/{fusion_mode}_raw{ch}/{ckpt_name}.pt") + if args.out == "": + out_path = Path(f"backup/{model}/{model_name}/{fusion_mode}_raw{ch}/{ckpt_name}.onnx") + + if fusion_mode != "stacked": + raise RuntimeError("Este exportador foi pensado para fusion_mode='stacked'.") + + 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) + + 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) + ) + + # Mesmo ajuste feito no treino: + # semantic usa o número real do 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()) + + print("==========================================") + print("Export SegFormer Multi-Head para ONNX") + print(f"Config : {config_path}") + print(f"Checkpoint : {checkpoint_path}") + print(f"Output ONNX : {out_path}") + print(f"Backbone : {backbone}") + print(f"Resolution : {W}x{H}") + 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"Resize output: {args.resize_to_input}") + print(f"Include norm : {args.include_norm}") + print(f"Postprocess : {args.postprocess}") + print("==========================================") + + 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 não disponível. Exportando em CPU.") + + model = train_mod.build_model( + backbone=backbone, + channels=channels, + heads_config=heads_config, + semantic_id2label=semantic_id2label, + semantic_label2id=semantic_label2id, + ) + + print("[CKPT] Carregando checkpoint...") + ckpt = torch.load( + str(checkpoint_path), + map_location="cpu", + weights_only=False, + ) + + if "model" not in ckpt: + raise RuntimeError( + "Checkpoint não contém a chave 'model'. " + "Confirme se é um checkpoint salvo pelo script de treino." + ) + + model.load_state_dict(ckpt["model"], strict=True) + model.to(device) + model.eval() + + norm_stats_path = None + norm_mean = None + norm_std = None + norm_channels = [] + + if args.include_norm: + norm_stats_path, norm_mean, norm_std, norm_channels = load_selected_norm_stats( + config=config, + config_path=config_path, + input_channel_indices=input_channel_indices, + ) + + print(f"[NORM] Embutindo normalização no ONNX: {norm_stats_path}") + print(f"[NORM] channels={norm_channels if norm_channels else input_channel_names}") + print(f"[NORM] mean={norm_mean}") + print(f"[NORM] std ={norm_std}") + + wrapper = MultiHeadOnnxWrapper( + model=model, + output_heads=output_heads, + resize_to_input=args.resize_to_input, + include_norm=args.include_norm, + norm_mean=norm_mean, + norm_std=norm_std, + postprocess=args.postprocess, + ) + wrapper.to(device) + wrapper.eval() + + dummy_input = torch.randn( + 1, + channels, + int(H), + int(W), + dtype=torch.float32, + device=device, + ) + + input_names = ["pixel_values"] + if args.postprocess in ("argmax_lowres", "argmax_fullres"): + output_names = [f"{name}_mask" for name in output_heads] + else: + output_names = [f"{name}_logits" for name in output_heads] + + dynamic_axes = None + if args.dynamic_batch: + dynamic_axes = { + "pixel_values": {0: "batch"}, + } + for out_name in output_names: + dynamic_axes[out_name] = {0: "batch"} + + print("[CHECK] Rodando forward PyTorch antes do export...") + with torch.no_grad(): + y = wrapper(dummy_input) + + for name, tensor in zip(output_names, y): + print(f" {name}: shape={tuple(tensor.shape)} dtype={tensor.dtype}") + + print("[EXPORT] Exportando ONNX...") + with torch.no_grad(): + torch.onnx.export( + wrapper, + dummy_input, + str(out_path), + export_params=True, + opset_version=int(args.opset), + do_constant_folding=True, + input_names=input_names, + output_names=output_names, + dynamic_axes=dynamic_axes, + ) + + print(f"[OK] ONNX salvo em: {out_path}") + + meta_path = out_path.with_suffix(".export_meta.json") + save_export_metadata(meta_path, { + "config": str(config_path), + "checkpoint": str(checkpoint_path), + "onnx": str(out_path), + "backbone": backbone, + "input_shape": [1, channels, int(H), int(W)], + "input_channel_names": input_channel_names, + "input_channel_indices": input_channel_indices, + "heads": output_heads, + "output_names": output_names, + "resize_to_input": bool(args.resize_to_input), + "opset": int(args.opset), + "dynamic_batch": bool(args.dynamic_batch), + "semantic_id2label": semantic_id2label, + "heads_config": heads_config, + "checkpoint_epoch": ckpt.get("epoch", None), + "checkpoint_best": ckpt.get("best", None), + "include_norm": bool(args.include_norm), + "norm_stats_path": str(norm_stats_path) if norm_stats_path is not None else None, + "norm_channels": norm_channels if norm_channels else input_channel_names, + "norm_mean": norm_mean, + "norm_std": norm_std, + "postprocess": str(args.postprocess), + }) + + print(f"[OK] Metadata salvo em: {meta_path}") + + # Verificação opcional do pacote onnx, se instalado. + try: + import onnx + print("[ONNX] Verificando grafo com onnx.checker...") + onnx_model = onnx.load(str(out_path)) + onnx.checker.check_model(onnx_model) + print("[OK] onnx.checker passou.") + except ImportError: + print("[WARN] Pacote onnx não instalado. Pulei onnx.checker.") + print(" Instale com: pip install onnx") + except Exception as e: + print(f"[WARN] onnx.checker encontrou problema: {e}") + + print("\nExport finalizado.") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/Python/OAK/datasets/oak-fcc-3/_11_validate_onnx.py b/Python/OAK/datasets/oak-fcc-3/_11_validate_onnx.py new file mode 100644 index 000000000..102f1ac2e --- /dev/null +++ b/Python/OAK/datasets/oak-fcc-3/_11_validate_onnx.py @@ -0,0 +1,1172 @@ +#!/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 + +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, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F + + +DEFAULT_CHANNEL_ORDER = ["R", "G", "B", "RE", "NIR"] + + +# ============================================================ +# 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], + channels: int, + channel_indices: List[int], + channel_names: List[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}") + print("[NORM] Sem norm_stats. Usando tensor 0..1 sem padronização.") + return None, None, None + + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + max_idx = max(channel_indices) + + if len(mean) <= max_idx or len(std) <= max_idx: + raise RuntimeError( + f"norm_stats incompatível: precisa índices={channel_indices}, " + f"mean={len(mean)} std={len(std)}" + ) + + mean_sel = [float(mean[i]) for i in channel_indices] + std_sel = [float(std[i]) for i in channel_indices] + + if names: + names_sel = [names[i] for i in channel_indices] + else: + names_sel = channel_names + + print(f"[NORM] usando {path}") + print(f"[NORM] channels={names_sel}") + 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, channels: int, channel_indices: List[int]) -> 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 arr.shape[0] in (3, 4, 5): + chw = arr + elif arr.shape[-1] in (3, 4, 5): + chw = np.transpose(arr, (2, 0, 1)) + else: + raise RuntimeError(f"Tensor com layout inesperado: {path} shape={arr.shape}") + + max_idx = max(channel_indices) + if chw.shape[0] <= max_idx: + raise RuntimeError( + f"Tensor {path} tem {chw.shape[0]} canais, " + f"mas precisa acessar índice {max_idx}." + ) + + chw = chw[channel_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=0.0) + return np.clip(chw, 0.0, 1.0).astype(np.float32) + + +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, Path.cwd()) + + W, H = config.get("resolucao", [1024, 640]) + dataset_path = config_dir / "dataset" + + candidates = [ + dataset_path / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json", + save_dir / "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 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}/{fusion_mode}_raw{channels}/{ckpt_name}.pt + backup/{modelo}/{model_name}/{fusion_mode}_raw{channels}/{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") + fusion_mode = config.get("fusion_mode", "stacked") + + # Melhor usar o channels real computado pelo script, + # porque ele vem do input_channel_names já resolvido. + ch = int(channels) + + ckpt_name = config.get("ckpt_test", "best_score") + + base_dir = config_dir / "backup" / model / model_name / f"{fusion_mode}_raw{ch}" + + if args.checkpoint: + checkpoint_path = resolve_path(args.checkpoint, Path.cwd()) + else: + checkpoint_path = base_dir / f"{ckpt_name}.pt" + + if args.onnx: + onnx_path = resolve_path(args.onnx, Path.cwd()) + else: + onnx_path = base_dir / f"{ckpt_name}.onnx" + + 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], + resize_to_input: bool = False, + output_kind: str = "logits", + ): + super().__init__() + self.model = model + self.output_heads = list(output_heads) + self.resize_to_input = bool(resize_to_input) + self.output_kind = str(output_kind).lower() + + if self.output_kind not in ("logits", "mask"): + raise RuntimeError(f"output_kind inválido: {self.output_kind}") + + 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.resize_to_input or self.output_kind == "mask": + if logits.shape[-2:] != input_hw: + logits = F.interpolate( + logits, + size=input_hw, + mode="bilinear", + align_corners=False, + ) + + if self.output_kind == "mask": + 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="Use quando o ONNX já inclui normalização interna. Nesse caso o ONNX recebe tensor 0..1, não tensor normalizado.", + ) + + parser.add_argument( + "--onnx_output_kind", + default="logits", + choices=["logits", "mask"], + help="Tipo de saída do ONNX: logits para modelo cru, mask para ONNX com argmax/postprocess embutido.", + ) + + 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()) + train_script_path = resolve_path(args.train_script, Path.cwd()) + labelmap_path = resolve_path(args.labelmap, Path.cwd()) + + if config_path is None or not config_path.is_file(): + raise FileNotFoundError(f"Config não encontrado: {config_path}") + 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}") + + config_dir = config_path.parent + config = load_json(config_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, + ) + + 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()) + + save_dir = ( + config_dir + / "backup" + / config.get("modelo", "segformer_b1") + / config.get("model_name", "test") + / f"{config.get('fusion_mode', 'stacked')}_raw{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, + channels=channels, + channel_indices=input_channel_indices, + 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"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"Compare HxW : {args.compare_at_input_size}") + print("==========================================") + + print("[MODEL] Montando PyTorch...") + model = train_mod.build_model( + backbone=backbone, + channels=channels, + heads_config=heads_config, + semantic_id2label=semantic_id2label, + semantic_label2id=semantic_label2id, + ) + + ckpt = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + if "model" not in ckpt: + raise RuntimeError("Checkpoint não contém chave 'model'.") + + model.load_state_dict(ckpt["model"], strict=True) + model.to(device) + model.eval() + + torch_wrapper = MultiHeadTorchWrapper( + model=model, + output_heads=output_heads, + resize_to_input=args.resize_torch_to_input or args.onnx_output_kind == "mask", + output_kind=args.onnx_output_kind, + ).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_name = onnx_session.get_inputs()[0].name + print(f"[ONNX] input name: {onnx_input_name}") + 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): + chw01 = load_tensor( + tensor_path, + channels=channels, + channel_indices=input_channel_indices, + ) + + # 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 args.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 args.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), + "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, + "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, Path.cwd()) + 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() \ No newline at end of file diff --git a/Python/OAK/datasets/oak-fcc-3/_12_benchmark_onnx.py b/Python/OAK/datasets/oak-fcc-3/_12_benchmark_onnx.py new file mode 100644 index 000000000..30383f3e0 --- /dev/null +++ b/Python/OAK/datasets/oak-fcc-3/_12_benchmark_onnx.py @@ -0,0 +1,946 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- + +""" +_12_benchmark_onnx.py + +Benchmark PyTorch vs ONNX Runtime para SegFormer OAK-FCC-3 Multi-Head. + +Mede: + - PyTorch FP32 + - PyTorch AMP/FP16 + - ONNX Runtime CUDA ou CPU + +Exemplos: + +Benchmark ONNX cru 160x256: + +python _12_benchmark_onnx.py --config config.json --max_samples 50 --warmup 10 --repeat 5 --device cuda --onnx_provider cuda + +Benchmark ONNX resized 640x1024: + +python _12_benchmark_onnx.py --config config.json --max_samples 50 --warmup 10 --repeat 5 --device cuda --onnx_provider cuda + + +TensorRT +python _12_benchmark_onnx.py --config config.json --max_samples 50 --warmup 10 --repeat 5 --device cuda --onnx_provider tensorrt --skip_torch_fp32 --skip_torch_amp + +""" + +from __future__ import annotations + +import gc +import csv +import json +import time +import argparse +import importlib.util +from pathlib import Path +from typing import Optional, List, Dict, Tuple + +import cv2 +import numpy as np +import torch +import torch.nn as nn + + +# ============================================================ +# 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 save_csv(path: str | Path, rows: List[dict]): + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + + if not rows: + return + + keys = list(rows[0].keys()) + + with path.open("w", newline="", encoding="utf-8") as f: + w = csv.DictWriter(f, fieldnames=keys) + w.writeheader() + w.writerows(rows) + + +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 synchronize_if_cuda(device: torch.device): + if device.type == "cuda": + torch.cuda.synchronize() + + +def clear_cuda(): + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize() + + +def percentile(values: List[float], p: float) -> float: + if not values: + return 0.0 + return float(np.percentile(np.asarray(values, dtype=np.float64), p)) + + +def summarize_times(times_ms: List[float]) -> dict: + arr = np.asarray(times_ms, dtype=np.float64) + + if arr.size == 0: + return { + "n": 0, + "mean_ms": 0.0, + "median_ms": 0.0, + "min_ms": 0.0, + "max_ms": 0.0, + "p95_ms": 0.0, + "p99_ms": 0.0, + "fps_mean": 0.0, + "fps_p95_latency": 0.0, + } + + mean_ms = float(arr.mean()) + p95_ms = float(np.percentile(arr, 95)) + p99_ms = float(np.percentile(arr, 99)) + + return { + "n": int(arr.size), + "mean_ms": mean_ms, + "median_ms": float(np.median(arr)), + "min_ms": float(arr.min()), + "max_ms": float(arr.max()), + "p95_ms": p95_ms, + "p99_ms": p99_ms, + "fps_mean": float(1000.0 / mean_ms) if mean_ms > 0 else 0.0, + "fps_p95_latency": float(1000.0 / p95_ms) if p95_ms > 0 else 0.0, + } + + +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}/{fusion_mode}_raw{channels}/{ckpt_name}.pt + backup/{modelo}/{model_name}/{fusion_mode}_raw{channels}/{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") + fusion_mode = config.get("fusion_mode", "stacked") + ckpt_name = config.get("ckpt_test", "best_score") + + # Usa o número real de canais selecionados, + # não necessariamente config["channels"]. + ch = int(channels) + + base_dir = config_dir / "backup" / model / model_name / f"{fusion_mode}_raw{ch}" + + if args.checkpoint: + checkpoint_path = resolve_path(args.checkpoint, Path.cwd()) + else: + checkpoint_path = base_dir / f"{ckpt_name}.pt" + + if args.onnx: + onnx_path = resolve_path(args.onnx, Path.cwd()) + else: + onnx_path = base_dir / f"{ckpt_name}.onnx" + + 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) + + +# ============================================================ +# Dataset / normalização +# ============================================================ + +def collect_tensor_samples(root: Path, max_samples: int = 50, 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 load_tensor(path: Path, channels: int, channel_indices: List[int]) -> 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 arr.shape[0] in (3, 4, 5): + chw = arr + elif arr.shape[-1] in (3, 4, 5): + chw = np.transpose(arr, (2, 0, 1)) + else: + raise RuntimeError(f"Tensor com layout inesperado: {path} shape={arr.shape}") + + max_idx = max(channel_indices) + if chw.shape[0] <= max_idx: + raise RuntimeError( + f"Tensor {path} tem {chw.shape[0]} canais, " + f"mas precisa acessar índice {max_idx}." + ) + + chw = chw[channel_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=0.0) + return np.clip(chw, 0.0, 1.0).astype(np.float32) + + +def find_norm_stats(config: dict, config_dir: Path, save_dir: Path, explicit: Optional[str]) -> Optional[Path]: + if explicit: + return resolve_path(explicit, Path.cwd()) + + W, H = config.get("resolucao", [1024, 640]) + dataset_path = config_dir / "dataset" + + candidates = [ + dataset_path / f"{int(W)}x{int(H)}" / "group" / "norm_stats.json", + save_dir / "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 load_norm_stats( + path: Optional[Path], + channel_indices: List[int], + channel_names: List[str], +) -> Tuple[Optional[List[float]], Optional[List[float]], Optional[str]]: + if path is None or not path.is_file(): + print("[NORM] Sem norm_stats. Usando tensor 0..1 sem padronização.") + return None, None, None + + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = js.get("channels", []) + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + + max_idx = max(channel_indices) + if len(mean) <= max_idx or len(std) <= max_idx: + raise RuntimeError( + f"norm_stats incompatível: precisa índices={channel_indices}, " + f"mean={len(mean)} std={len(std)}" + ) + + mean_sel = [float(mean[i]) for i in channel_indices] + std_sel = [float(std[i]) for i in channel_indices] + + if names: + names_sel = [names[i] for i in channel_indices] + else: + names_sel = channel_names + + print(f"[NORM] usando {path}") + print(f"[NORM] channels={names_sel}") + 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_inputs_as_numpy( + samples: List[Path], + channels: int, + channel_indices: List[int], + mean: Optional[List[float]], + std: Optional[List[float]], + target_hw: Tuple[int, int], + normalize_input: bool = True, +) -> List[np.ndarray]: + H, W = target_hw + xs = [] + + for p in samples: + chw01 = load_tensor( + p, + channels=channels, + channel_indices=channel_indices, + ) + + 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) + + if normalize_input: + chw = normalize_numpy_chw(chw01, mean=mean, std=std) + else: + chw = chw01.astype(np.float32, copy=False) + + x = np.expand_dims(chw, axis=0).astype(np.float32) + xs.append(x) + + return xs + + +# ============================================================ +# PyTorch +# ============================================================ + +class TorchTupleWrapper(nn.Module): + def __init__(self, model: nn.Module, output_heads: List[str]): + super().__init__() + self.model = model + self.output_heads = list(output_heads) + + def forward(self, pixel_values: torch.Tensor): + outputs = self.model(pixel_values=pixel_values) + return tuple(outputs[h] for h in self.output_heads) + + +@torch.inference_mode() +def benchmark_torch( + model: nn.Module, + inputs_np: List[np.ndarray], + device: torch.device, + warmup: int, + repeat: int, + amp: bool, + label: str, +) -> Tuple[dict, List[dict]]: + model.eval() + + times = [] + rows = [] + + # Precarrega tensors na GPU para medir só inferência do modelo. + inputs_t = [ + torch.from_numpy(x).to(device, non_blocking=True) + for x in inputs_np + ] + + if device.type == "cuda": + torch.cuda.synchronize() + + print(f"\n[BENCH] {label} | warmup={warmup} repeat={repeat}") + + # Warmup + for i in range(max(0, warmup)): + x = inputs_t[i % len(inputs_t)] + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=amp and device.type == "cuda"): + _ = model(x) + + synchronize_if_cuda(device) + + # Medição + total_iter = len(inputs_t) * max(1, repeat) + idx = 0 + + for r in range(max(1, repeat)): + for sample_idx, x in enumerate(inputs_t): + synchronize_if_cuda(device) + t0 = time.perf_counter() + + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=amp and device.type == "cuda"): + _ = model(x) + + synchronize_if_cuda(device) + dt_ms = (time.perf_counter() - t0) * 1000.0 + + times.append(dt_ms) + rows.append({ + "engine": label, + "repeat": r, + "sample_idx": sample_idx, + "iter_idx": idx, + "latency_ms": dt_ms, + }) + + idx += 1 + + if idx % 25 == 0 or idx == total_iter: + print(f" {idx:04d}/{total_iter:04d} | last={dt_ms:.2f}ms") + + summary = summarize_times(times) + summary["engine"] = label + + return summary, rows + + +# ============================================================ +# ONNX Runtime +# ============================================================ + +def create_onnx_session(onnx_path: Path, provider: str): + import os + + trt_home = os.environ.get("TRT_HOME", r"C:\dev\TensorRT-10.10.0.31") + + for dll_dir in [ + os.path.join(trt_home, "lib"), + os.path.join(trt_home, "bin"), + r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.4\bin", + ]: + 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}") + + try: + import onnxruntime as ort + except ImportError: + raise ImportError( + "onnxruntime não está instalado. Instale com:\n" + " pip install onnxruntime-gpu\n" + "ou CPU:\n" + " pip install onnxruntime" + ) + + available = ort.get_available_providers() + print(f"[ONNX] providers disponíveis: {available}") + + provider = provider.lower() + + 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, + + # FP16: o ponto principal do nosso teste. + "trt_fp16_enable": True, + + # Cache: evita rebuild do engine a cada execução. + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str(cache_dir), + + # Timing cache ajuda a acelerar builds futuros. + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(cache_dir), + + # Workspace. 4GB é razoável para RTX 3070, ajuste se faltar VRAM. + "trt_max_workspace_size": 4 * 1024 * 1024 * 1024, + } + + providers = [ + ("TensorrtExecutionProvider", trt_options), + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + else: + providers = [provider] + + # Checagem de disponibilidade, lidando com provider tuple. + 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}" + ) + + session = ort.InferenceSession( + str(onnx_path), + sess_options=sess_options, + providers=providers_ok, + ) + + print(f"[ONNX] usando providers: {session.get_providers()}") + + active_providers = session.get_providers() + + if provider == "tensorrt" and "TensorrtExecutionProvider" not in active_providers: + raise RuntimeError( + "TensorRTExecutionProvider foi solicitado, mas não ficou ativo. " + f"Providers ativos: {active_providers}. " + "Provável causa: TensorRT não instalado, DLLs fora do PATH, " + "ou versão incompatível com onnxruntime-gpu." + ) + + if provider == "cuda" and "CUDAExecutionProvider" not in active_providers: + raise RuntimeError( + "CUDAExecutionProvider foi solicitado, mas não ficou ativo. " + f"Providers ativos: {active_providers}." + ) + + return session + + +def benchmark_onnx( + session, + inputs_np: List[np.ndarray], + warmup: int, + repeat: int, + label: str, +) -> Tuple[dict, List[dict]]: + input_name = session.get_inputs()[0].name + + times = [] + rows = [] + + print(f"\n[BENCH] {label} | warmup={warmup} repeat={repeat}") + + # Warmup + for i in range(max(0, warmup)): + x = inputs_np[i % len(inputs_np)] + _ = session.run(None, {input_name: x}) + + # Medição + total_iter = len(inputs_np) * max(1, repeat) + idx = 0 + + for r in range(max(1, repeat)): + for sample_idx, x in enumerate(inputs_np): + t0 = time.perf_counter() + _ = session.run(None, {input_name: x}) + dt_ms = (time.perf_counter() - t0) * 1000.0 + + times.append(dt_ms) + rows.append({ + "engine": label, + "repeat": r, + "sample_idx": sample_idx, + "iter_idx": idx, + "latency_ms": dt_ms, + }) + + idx += 1 + + if idx % 25 == 0 or idx == total_iter: + print(f" {idx:04d}/{total_iter:04d} | last={dt_ms:.2f}ms") + + summary = summarize_times(times) + summary["engine"] = label + + return summary, rows + + +# ============================================================ +# 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=50) + parser.add_argument("--start_idx", type=int, default=0) + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--repeat", type=int, default=5) + + parser.add_argument("--device", default="cuda", choices=["cuda", "cpu"]) + parser.add_argument("--onnx_provider", default="cuda", choices=["cuda", "cpu", "tensorrt"]) + + parser.add_argument("--skip_torch_fp32", action="store_true") + parser.add_argument("--skip_torch_amp", action="store_true") + parser.add_argument("--skip_onnx", action="store_true") + + parser.add_argument( + "--onnx_has_norm", + action="store_true", + help="Use quando o ONNX já inclui normalização interna. Nesse caso o ONNX recebe tensor 0..1 cru.", + ) + + parser.add_argument("--out_dir", default=None) + + args = parser.parse_args() + + config_path = resolve_path(args.config, Path.cwd()) + train_script_path = resolve_path(args.train_script, Path.cwd()) + labelmap_path = resolve_path(args.labelmap, Path.cwd()) + + if config_path is None or not config_path.is_file(): + raise FileNotFoundError(f"Config não encontrado: {config_path}") + 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}") + + config_dir = config_path.parent + config = load_json(config_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, + ) + + 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()) + + save_dir = ( + config_dir + / "backup" + / config.get("modelo", "segformer_b1") + / config.get("model_name", "test") + / f"{config.get('fusion_mode', 'stacked')}_raw{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_indices=input_channel_indices, + 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, + ) + + if args.out_dir: + out_dir = resolve_path(args.out_dir, Path.cwd()) + else: + out_dir = onnx_path.parent / "benchmarks" + + assert out_dir is not None + out_dir.mkdir(parents=True, exist_ok=True) + + 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("Benchmark PyTorch vs ONNX") + print(f"Config : {config_path}") + print(f"Checkpoint : {checkpoint_path}") + print(f"ONNX : {onnx_path}") + print(f"Root : {root}") + print(f"Samples : {len(samples)}") + print(f"Warmup : {args.warmup}") + print(f"Repeat : {args.repeat}") + 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"ONNX has norm: {args.onnx_has_norm}") + print(f"Out dir : {out_dir}") + print("==========================================") + + if args.onnx_has_norm: + print("\n[DATA] Carregando inputs 0..1 crus na RAM...") + else: + print("\n[DATA] Carregando inputs normalizados na RAM...") + inputs_np = load_inputs_as_numpy( + samples=samples, + channels=channels, + channel_indices=input_channel_indices, + mean=mean, + std=std, + target_hw=(H, W), + normalize_input=not args.onnx_has_norm, + ) + print(f"[DATA] Inputs carregados: {len(inputs_np)}") + + summaries = [] + all_rows = [] + + # ======================================================== + # PyTorch + # ======================================================== + need_torch = not args.skip_torch_fp32 or not args.skip_torch_amp + + if need_torch: + print("\n[MODEL] Montando PyTorch...") + model = train_mod.build_model( + backbone=backbone, + channels=channels, + heads_config=heads_config, + semantic_id2label=semantic_id2label, + semantic_label2id=semantic_label2id, + ) + + ckpt = torch.load(str(checkpoint_path), map_location="cpu", weights_only=False) + if "model" not in ckpt: + raise RuntimeError("Checkpoint não contém chave 'model'.") + + model.load_state_dict(ckpt["model"], strict=True) + model.to(device) + model.eval() + + torch_model = TorchTupleWrapper( + model=model, + output_heads=output_heads, + ).to(device) + torch_model.eval() + + clear_cuda() + + if not args.skip_torch_fp32: + summary, rows = benchmark_torch( + model=torch_model, + inputs_np=inputs_np, + device=device, + warmup=args.warmup, + repeat=args.repeat, + amp=False, + label="torch_fp32", + ) + summaries.append(summary) + all_rows.extend(rows) + + clear_cuda() + + if not args.skip_torch_amp: + summary, rows = benchmark_torch( + model=torch_model, + inputs_np=inputs_np, + device=device, + warmup=args.warmup, + repeat=args.repeat, + amp=True, + label="torch_amp_fp16", + ) + summaries.append(summary) + all_rows.extend(rows) + + del torch_model + del model + clear_cuda() + + # ======================================================== + # ONNX + # ======================================================== + if not args.skip_onnx: + print("\n[ONNX] Carregando sessão...") + session = create_onnx_session( + onnx_path=onnx_path, + provider=args.onnx_provider, + ) + + summary, rows = benchmark_onnx( + session=session, + inputs_np=inputs_np, + warmup=args.warmup, + repeat=args.repeat, + label=f"onnx_{args.onnx_provider}", + ) + summaries.append(summary) + all_rows.extend(rows) + + # ======================================================== + # Relatório + # ======================================================== + print("\n========== RESUMO ==========") + + for s in summaries: + print( + f"{s['engine']:<16} " + f"n={s['n']:<4} " + f"mean={s['mean_ms']:.3f}ms " + f"median={s['median_ms']:.3f}ms " + f"p95={s['p95_ms']:.3f}ms " + f"p99={s['p99_ms']:.3f}ms " + f"fps_mean={s['fps_mean']:.2f} " + f"fps_p95={s['fps_p95_latency']:.2f}" + ) + + base_name = f"{onnx_path.stem}_{args.onnx_provider}" + report_json = out_dir / f"{base_name}_benchmark_report.json" + report_csv = out_dir / f"{base_name}_benchmark_rows.csv" + + report = { + "config": str(config_path), + "checkpoint": str(checkpoint_path), + "ckpt_name": ckpt_name, + "onnx": str(onnx_path), + "onnx_has_norm": bool(args.onnx_has_norm), + "root": str(root), + "samples": len(samples), + "warmup": int(args.warmup), + "repeat": int(args.repeat), + "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, + "onnx_provider": args.onnx_provider, + "device": str(device), + "summaries": summaries, + } + + save_json(report_json, report) + save_csv(report_csv, all_rows) + + print(f"\n[OK] JSON salvo em: {report_json}") + print(f"[OK] CSV salvo em : {report_csv}") + print("\nBenchmark finalizado.") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/Python/OAK/datasets/oak-fcc-3/_5_augmentation.py b/Python/OAK/datasets/oak-fcc-3/_5_augmentation.py index a6d3cc6e2..37043ca86 100644 --- a/Python/OAK/datasets/oak-fcc-3/_5_augmentation.py +++ b/Python/OAK/datasets/oak-fcc-3/_5_augmentation.py @@ -2,439 +2,1620 @@ # -*- coding: utf-8 -*- """ -_5_augmentation_grouped_oak.py +_5_augmentation_raw_oak.py -Augmenta previews e máscaras por grupo na estrutura nova OAK-FCC-3. +Augmentation multiespectral para OAK-FCC-3 usando os .bin RAW das câmeras, +máscaras e metadados do dataset. -Entrada: +Objetivo: +- Ler amostras em dataset/original/group// +- Carregar RGB/RE/NIR a partir dos .bin RAW10 packed +- Aplicar augmentation geométrico sincronizado em todos os canais + máscara +- Aplicar augmentation radiométrico/físico nos RAWs +- Salvar nova amostra em dataset/augmented/group// +- Gerar preview RGB novo para inspeção visual -dataset/originals/group// - previews/ +Estrutura esperada de entrada: + +dataset/original/group// + bins/ masks/ metas/ - bins/ + previews/ # opcional -Saída: +Estrutura de saída: dataset/augmented/group// - previews/ + bins/ masks/ + metas/ + previews/ -Uso: +Exemplos: -python _5_augmentation_grouped_oak.py --copies 5 +python _5_augmentation_raw_oak.py --copies 3 --clear-dst -python _5_augmentation_grouped_oak.py ^ - --copies 5 ^ - --groups chao,chao_cana,chao_erva,chao_cana_erva - -python _5_augmentation_grouped_oak.py ^ - --copies 5 ^ - --src-root dataset/originals/group ^ - --dst-root dataset/augmented/group ^ +python _5_augmentation_raw_oak.py ^ + --copies 3 ^ + --groups chao,chao_cana,chao_erva,chao_cana_erva ^ --clear-dst + +python _5_augmentation_raw_oak.py ^ + --group-copies chao:1,chao_cana:3,chao_erva:3,chao_cana_erva:5 ^ + --clear-dst + +python _5_augmentation_raw_oak.py --dry-run + +Observações importantes: +- Este script NÃO augmenta preview PNG como fonte principal. +- O preview é somente consequência visual do RGB RAW augmentado. +- A geometria é sempre igual em RGB/RE/NIR/mask. +- Radiometria usa mesma base global + pequenas variações por banda/câmera. +- O script tenta ser tolerante a nomes de arquivo, mas espera que base da amostra + esteja preservada nos nomes dos bins/masks/metas. """ -import os -import json import argparse +import copy +import json +import math +import os import random import shutil +from dataclasses import dataclass, asdict from pathlib import Path +from typing import Dict, List, Optional, Tuple, Any import cv2 +import numpy as np from PIL import Image -import albumentations as A + +try: + from core.raw_processor_core import RawProcessorCore + _HAS_RAW_PROCESSOR_CORE = True +except Exception: + RawProcessorCore = None + _HAS_RAW_PROCESSOR_CORE = False + +_CORE_PREVIEW_CACHE = {} -# ====================== Configurações ====================== +# ============================================================ +# Configuração base +# ============================================================ -with open("config.json", "r", encoding="utf-8") as f: - config = json.load(f) +IMG_EXTS = (".png", ".jpg", ".jpeg") +META_EXTS = (".json",) +BIN_EXTS = (".bin",) -MODELO = config.get("camera", ".") +ROLE_ALIASES = { + "rgb": ["rgb", "color", "cor", "cam_a", "cama", "CAM_A"], + "re": ["re", "rededge", "red_edge", "red-edge", "cam_b", "camb", "CAM_B"], + "nir": ["nir", "ir", "infra", "infrared", "cam_c", "camc", "CAM_C"], +} -DATASET_BASE = os.path.join(MODELO, "dataset") +CAM_KEYS = ["CAM_A", "CAM_B", "CAM_C", "cam_a", "cam_b", "cam_c", "cam0", "cam1", "cam2"] -DEFAULT_SRC_ROOT = os.path.join(DATASET_BASE, "originals", "group") -DEFAULT_DST_ROOT = os.path.join(DATASET_BASE, "augmented", "group") - -IMG_EXTS = (".jpg", ".jpeg", ".png") -MSK_EXTS = (".png", ".jpg", ".jpeg") +DEFAULT_BIT_DEPTH = 10 +DEFAULT_BAYER_PATTERN = "BGGR" +DEFAULT_RAW_FORMAT = "RAW10_PACKED" -# ====================== Augmentation ====================== +# ============================================================ +# Data classes +# ============================================================ -train_tf = A.Compose( - [ - A.HorizontalFlip(p=0.5), - - A.ShiftScaleRotate( - shift_limit=0.01, - scale_limit=0.10, - rotate_limit=5, - border_mode=cv2.BORDER_REFLECT_101, - interpolation=cv2.INTER_LINEAR, - mask_interpolation=cv2.INTER_NEAREST, - p=0.30, - ), - - A.OneOf( - [ - A.RandomBrightnessContrast(0.2, 0.2, p=1.0), - A.HueSaturationValue( - hue_shift_limit=5, - sat_shift_limit=20, - val_shift_limit=15, - p=1.0, - ), - A.RandomGamma(gamma_limit=(90, 110), p=1.0), - ], - p=0.70, - ), - - A.OneOf( - [ - A.MotionBlur(blur_limit=3, p=1.0), - A.GaussianBlur(blur_limit=3, p=1.0), - ], - p=0.20, - ), - - A.OneOf( - [ - A.GaussNoise(var_limit=(5.0, 15.0), p=1.0), - A.ImageCompression(quality_lower=60, quality_upper=90, p=1.0), - ], - p=0.20, - ), - - A.RandomShadow(p=0.08), - A.RandomSunFlare(p=0.06), - - # Cuidado com ChannelShuffle em imagens agrícolas: - # pode bagunçar demais a semântica das cores. - # Mantive desligado por padrão. - # A.ChannelShuffle(p=0.03), - - A.CoarseDropout( - max_holes=6, - max_height=16, - max_width=16, - p=0.08, - ), - ] -) +@dataclass +class CameraBin: + role: str + cam_key: str + path: Path + width: int + height: int + bit_depth: int = DEFAULT_BIT_DEPTH + raw_format: str = DEFAULT_RAW_FORMAT + bayer_pattern: str = DEFAULT_BAYER_PATTERN + data: Optional[np.ndarray] = None -# ====================== Utilitários ====================== +@dataclass +class AugParams: + seed: int -def garantir_dir(p): - os.makedirs(p, exist_ok=True) + # Geometria, igual para tudo + flip_h: bool + shift_x_frac: float + shift_y_frac: float + scale: float + rotate_deg: float + perspective: bool + perspective_strength: float + + # Radiometria base + exposure_mult: float + gamma: float + + # Ganhos por papel espectral + rgb_channel_gain: Tuple[float, float, float] + re_gain: float + nir_gain: float + + # Efeitos físicos leves + shadow_enabled: bool + shadow_strength: float + shadow_angle_deg: float + highlight_enabled: bool + highlight_strength: float + noise_enabled: bool + noise_sigma_dn: float + blur_enabled: bool + blur_kernel: int -def maybe_clear_dir(p: str): - if os.path.isdir(p): - shutil.rmtree(p) - garantir_dir(p) +# ============================================================ +# Utilitários gerais +# ============================================================ + +def read_config_model() -> str: + """Tenta ler config.json para manter compatibilidade com seus scripts atuais.""" + cfg = Path("config.json") + if not cfg.exists(): + return "." + + try: + with cfg.open("r", encoding="utf-8") as f: + config = json.load(f) + return str(config.get("camera", ".")) + except Exception: + return "." + + +def read_config_value(key: str, default=None): + cfg = Path("config.json") + if not cfg.exists(): + return default + + try: + with cfg.open("r", encoding="utf-8") as f: + config = json.load(f) + return config.get(key, default) + except Exception: + return default + + +def ensure_dir(path: Path) -> None: + path.mkdir(parents=True, exist_ok=True) + + +def clear_dir(path: Path) -> None: + if path.exists(): + shutil.rmtree(path) + ensure_dir(path) def normalizar_base(stem: str) -> str: - sufixos = [ - "_rgb", "_RGB", "_Rgb", - "_image", "_img", "_frame", - "_preview", "_previews", - "_mask", "_masks", - "_seg", "_SEG", "_segment", "_segmentacao", "_Segmentacao", + """Remove sufixos comuns para casar preview/mask/meta/bin.""" + suffixes = [ + "_rgb", "_RGB", "_Rgb", "_color", "_COLOR", + "_re", "_RE", "_rededge", "_red_edge", "_nir", "_NIR", + "_cam_a", "_cam_b", "_cam_c", "_CAM_A", "_CAM_B", "_CAM_C", + "_cam0", "_cam1", "_cam2", "_CAM0", "_CAM1", "_CAM2", + "_image", "_img", "_frame", "_preview", "_previews", + "_mask", "_masks", "_seg", "_SEG", "_segment", "_segmentacao", + "_meta", "_metadata", ] out = stem - mudou = True - - while mudou: - mudou = False - for sfx in sufixos: + changed = True + while changed: + changed = False + for sfx in suffixes: if out.endswith(sfx): out = out[: -len(sfx)] - mudou = True + changed = True break - return out -def map_files_by_base(folder, exts): - by_base = {} +def list_groups(src_root: Path) -> List[str]: + if not src_root.exists(): + return [] - if not os.path.isdir(folder): + groups = [] + for p in sorted(src_root.iterdir()): + if not p.is_dir(): + continue + if (p / "bins").is_dir() and (p / "masks").is_dir() and (p / "metas").is_dir(): + groups.append(p.name) + return groups + + +def map_files_by_base(folder: Path, exts: Tuple[str, ...]) -> Dict[str, Path]: + by_base: Dict[str, Path] = {} + if not folder.is_dir(): return by_base - prioridade = { + priority = { + ".json": 0, ".png": 0, + ".bin": 0, ".jpg": 1, ".jpeg": 2, } - for fname in os.listdir(folder): - lower = fname.lower() - - if not lower.endswith(exts): + for p in sorted(folder.iterdir()): + if not p.is_file(): + continue + ext = p.suffix.lower() + if ext not in exts: continue - stem, ext = os.path.splitext(fname) - base = normalizar_base(stem) - cand = os.path.join(folder, fname) - + base = normalizar_base(p.stem) if base not in by_base: - by_base[base] = cand + by_base[base] = p else: - cur_ext = os.path.splitext(by_base[base])[1].lower() - if prioridade.get(ext.lower(), 99) < prioridade.get(cur_ext, 99): - by_base[base] = cand + cur_ext = by_base[base].suffix.lower() + if priority.get(ext, 99) < priority.get(cur_ext, 99): + by_base[base] = p return by_base -def list_groups(src_root): - if not os.path.isdir(src_root): - return [] - - grupos = [] - - for name in sorted(os.listdir(src_root)): - gdir = os.path.join(src_root, name) - - if not os.path.isdir(gdir): - continue - - previews = os.path.join(gdir, "previews") - masks = os.path.join(gdir, "masks") - - if os.path.isdir(previews) and os.path.isdir(masks): - grupos.append(name) - - return grupos +def load_json(path: Path) -> Dict[str, Any]: + with path.open("r", encoding="utf-8") as f: + return json.load(f) -def load_rgb(path): - im = cv2.imread(path, cv2.IMREAD_COLOR) +def save_json(path: Path, data: Dict[str, Any]) -> None: + ensure_dir(path.parent) + with path.open("w", encoding="utf-8") as f: + json.dump(data, f, ensure_ascii=False, indent=2) + +def load_mask(path: Path) -> np.ndarray: + im = cv2.imread(str(path), cv2.IMREAD_UNCHANGED) if im is None: raise FileNotFoundError(path) - return cv2.cvtColor(im, cv2.COLOR_BGR2RGB) + if im.ndim == 2: + return im + + if im.shape[2] == 4: + im = cv2.cvtColor(im, cv2.COLOR_BGRA2RGBA) + else: + im = cv2.cvtColor(im, cv2.COLOR_BGR2RGB) + return im -def save_rgb(path, arr_rgb): - garantir_dir(os.path.dirname(path)) - Image.fromarray(arr_rgb).save(path) +def save_mask(path: Path, mask: np.ndarray) -> None: + ensure_dir(path.parent) + if mask.ndim == 2: + Image.fromarray(mask).save(path) + else: + Image.fromarray(mask).save(path) -def ensure_aug_dirs(dst_root, group_name): - preview_out = os.path.join(dst_root, group_name, "previews") - mask_out = os.path.join(dst_root, group_name, "masks") +# ============================================================ +# RAW10 pack/unpack +# ============================================================ - garantir_dir(preview_out) - garantir_dir(mask_out) - - return preview_out, mask_out - - -def validate_mask_colors_before_after(mask_before, mask_after, name="mask"): +def unpack_raw10_packed(buf: bytes, width: int, height: int) -> np.ndarray: """ - Só diagnóstico leve: máscaras não devem ganhar milhares de cores. - Em segmentação por cor, interpolação errada cria sujeira cromática. - """ - before_colors = len(set(map(tuple, mask_before.reshape(-1, 3)))) - after_colors = len(set(map(tuple, mask_after.reshape(-1, 3)))) + Desempacota MIPI RAW10 packed no padrão mais comum usado por câmeras/DepthAI: - if after_colors > max(before_colors * 3, 32): - print( - f"[WARN] {name}: muitas cores após aug. " - f"antes={before_colors}, depois={after_colors}. " - f"Confira se mask_interpolation está NEAREST." + Para cada grupo de 5 bytes: + byte0 = bits [9:2] do pixel 0 + byte1 = bits [9:2] do pixel 1 + byte2 = bits [9:2] do pixel 2 + byte3 = bits [9:2] do pixel 3 + byte4 = bits [1:0] dos 4 pixels, empilhados em pares de bits + + Retorna uint16 com valores 0..1023. + + Importante: + - A versão anterior tratava byte0..byte3 como bits baixos e byte4 como bits altos. + Isso embaralha o RAW e gera exatamente aquele padrão de ruído/colorido sem imagem. + """ + expected_pixels = width * height + arr = np.frombuffer(buf, dtype=np.uint8) + + expected_bytes = (expected_pixels * 10 + 7) // 8 + if arr.size < expected_bytes: + raise ValueError( + f"RAW10 menor que esperado: bytes={arr.size}, esperado>={expected_bytes}, " + f"shape={width}x{height}" ) + arr = arr[:expected_bytes] + groups = arr.size // 5 + main = arr[: groups * 5].reshape(-1, 5).astype(np.uint16) -# ====================== Augmentação ====================== + p0 = (main[:, 0] << 2) | ((main[:, 4] >> 0) & 0x03) + p1 = (main[:, 1] << 2) | ((main[:, 4] >> 2) & 0x03) + p2 = (main[:, 2] << 2) | ((main[:, 4] >> 4) & 0x03) + p3 = (main[:, 3] << 2) | ((main[:, 4] >> 6) & 0x03) -def augment_pair( - preview_path, - mask_path, - preview_out_dir, - mask_out_dir, - copies, -): - base_preview, preview_ext = os.path.splitext(os.path.basename(preview_path)) - _, mask_ext = os.path.splitext(os.path.basename(mask_path)) + out = np.empty(groups * 4, dtype=np.uint16) + out[0::4] = p0 + out[1::4] = p1 + out[2::4] = p2 + out[3::4] = p3 - base = normalizar_base(base_preview) + if out.size < expected_pixels: + raise ValueError(f"RAW10 gerou pixels insuficientes: {out.size} < {expected_pixels}") - img = load_rgb(preview_path) - mask = load_rgb(mask_path) - - gen = 0 - - for i in range(copies): - aug = train_tf(image=img, mask=mask) - - img_aug = aug["image"] - mask_aug = aug["mask"] - - validate_mask_colors_before_after(mask, mask_aug, name=f"{base}_aug_{i:02d}") - - new_base = f"{base}_aug_{i:02d}" - - out_preview = os.path.join(preview_out_dir, f"{new_base}{preview_ext}") - out_mask = os.path.join(mask_out_dir, f"{new_base}{mask_ext}") - - save_rgb(out_preview, img_aug) - save_rgb(out_mask, mask_aug) - - gen += 1 - - return gen + return out[:expected_pixels].reshape(height, width) -# ====================== Processamento ====================== +def pack_raw10_packed(raw: np.ndarray) -> bytes: + """ + Empacota uint16 0..1023 para MIPI RAW10 packed no mesmo padrão do unpack: -def process_group( - src_root, - dst_root, - group_name, - copies, - limit=None, - seed=42, -): - group_dir = os.path.join(src_root, group_name) + byte0..byte3 = bits altos [9:2] + byte4 = bits baixos [1:0] de p0,p1,p2,p3 + """ + flat = np.asarray(raw, dtype=np.uint16).reshape(-1) + flat = np.clip(flat, 0, 1023).astype(np.uint16) - preview_dir = os.path.join(group_dir, "previews") - mask_dir = os.path.join(group_dir, "masks") + pad = (-flat.size) % 4 + if pad: + flat = np.pad(flat, (0, pad), mode="edge") - if not os.path.isdir(preview_dir) or not os.path.isdir(mask_dir): - print(f"[WARN] Grupo '{group_name}' inválido, sem previews/masks. Pulando.") - return 0 + p = flat.reshape(-1, 4) + out = np.empty((p.shape[0], 5), dtype=np.uint8) - preview_map = map_files_by_base(preview_dir, IMG_EXTS) - mask_map = map_files_by_base(mask_dir, MSK_EXTS) + out[:, 0] = ((p[:, 0] >> 2) & 0xFF).astype(np.uint8) + out[:, 1] = ((p[:, 1] >> 2) & 0xFF).astype(np.uint8) + out[:, 2] = ((p[:, 2] >> 2) & 0xFF).astype(np.uint8) + out[:, 3] = ((p[:, 3] >> 2) & 0xFF).astype(np.uint8) + out[:, 4] = ( + ((p[:, 0] & 0x03) << 0) + | ((p[:, 1] & 0x03) << 2) + | ((p[:, 2] & 0x03) << 4) + | ((p[:, 3] & 0x03) << 6) + ).astype(np.uint8) - bases = sorted(set(preview_map.keys()) & set(mask_map.keys())) + return out.reshape(-1).tobytes() - if limit is not None and limit > 0 and limit < len(bases): - rng = random.Random(seed) - bases = sorted(rng.sample(bases, limit)) - print(f"[INFO] [{group_name}] Limit aplicado: {limit} amostras originais selecionadas.") - preview_out_dir, mask_out_dir = ensure_aug_dirs(dst_root, group_name) +def pack_raw10_packed_array(raw: np.ndarray) -> np.ndarray: + """ + RAW16 0..1023 -> ndarray uint8 RAW10 packed shape=(height, packed_width). + Esse é o formato esperado pelo RawProcessorCore.decode_stream_cameras. + """ + h, w = raw.shape[:2] + packed = np.frombuffer(pack_raw10_packed(raw), dtype=np.uint8) + packed_w = int(math.ceil(int(w) * 10 / 8)) + return packed.reshape(int(h), packed_w) - count = 0 - sem_mask = 0 - erros = 0 - for base in bases: - preview_path = preview_map.get(base) - mask_path = mask_map.get(base) +def read_raw_bin(path: Path, width: int, height: int, raw_format: str) -> np.ndarray: + raw_format_u = (raw_format or "").upper() + buf = path.read_bytes() - if not preview_path or not mask_path: - sem_mask += 1 - print(f"[WARN] [{group_name}] par incompleto para base={base}. Pulando.") + if "RAW10" in raw_format_u or "PACKED" in raw_format_u: + return unpack_raw10_packed(buf, width, height) + + # Fallback: RAW16 little-endian com valores possivelmente 0..1023. + arr = np.frombuffer(buf, dtype=np.uint16) + expected = width * height + if arr.size < expected: + raise ValueError(f"RAW16 menor que esperado em {path}: {arr.size} < {expected}") + return arr[:expected].reshape(height, width) + + +def write_raw_bin(path: Path, raw: np.ndarray, raw_format: str) -> None: + ensure_dir(path.parent) + raw_format_u = (raw_format or "").upper() + + raw10 = np.clip(np.rint(raw), 0, 1023).astype(np.uint16) + + if "RAW10" in raw_format_u or "PACKED" in raw_format_u: + path.write_bytes(pack_raw10_packed(raw10)) + else: + path.write_bytes(raw10.astype(np.uint16).tobytes()) + + +# ============================================================ +# Metadados e descoberta de bins +# ============================================================ + +def _text_has_any(text: str, needles: List[str]) -> bool: + t = text.lower() + return any(n.lower() in t for n in needles) + + +def role_from_text(text: str) -> Optional[str]: + for role, aliases in ROLE_ALIASES.items(): + if _text_has_any(text, aliases): + return role + return None + + +def normalize_role(role: Optional[str], cam_key: Optional[str] = None) -> Optional[str]: + if role: + r = str(role).lower().strip() + if r in ("rgb", "color", "cor"): + return "rgb" + if r in ("re", "rededge", "red_edge", "red-edge"): + return "re" + if r in ("nir", "ir", "infrared", "infra"): + return "nir" + + if cam_key: + ck = str(cam_key).upper() + # Convenção atual mais provável: CAM_A=RGB, CAM_B=RE, CAM_C=NIR. + if ck == "CAM_A": + return "rgb" + if ck == "CAM_B": + return "re" + if ck == "CAM_C": + return "nir" + + return None + + +def extract_camera_info(meta: Dict[str, Any]) -> Dict[str, Dict[str, Any]]: + """ + Tenta extrair camera_info de vários formatos possíveis. + Retorna dict cam_key -> info. + """ + candidates = [] + + for key in ["camera_info", "cameras", "camera_infos", "payload_sources_info"]: + if isinstance(meta.get(key), dict): + candidates.append(meta[key]) + + # Algumas versões guardam dentro de meta["capture"] ou meta["meta"] + for parent_key in ["capture", "meta", "frame_meta"]: + parent = meta.get(parent_key) + if isinstance(parent, dict): + for key in ["camera_info", "cameras", "camera_infos"]: + if isinstance(parent.get(key), dict): + candidates.append(parent[key]) + + if candidates: + cam_info = candidates[0] + out = {} + for k, v in cam_info.items(): + if isinstance(v, dict): + out[str(k)] = v + return out + + # Fallback mínimo caso o meta já tenha campos diretos. + out = {} + for ck in CAM_KEYS: + if isinstance(meta.get(ck), dict): + out[ck] = meta[ck] + return out + + +def get_int_any(d: Dict[str, Any], keys: List[str], default: Optional[int] = None) -> Optional[int]: + for k in keys: + if k in d and d[k] is not None: + try: + return int(d[k]) + except Exception: + pass + return default + + +def get_str_any(d: Dict[str, Any], keys: List[str], default: str = "") -> str: + for k in keys: + if k in d and d[k] is not None: + return str(d[k]) + return default + + +def find_bin_for_camera(bins_dir: Path, base: str, role: str, cam_key: str) -> Optional[Path]: + """ + Localiza o bin de uma câmera por base + role/cam_key. + É tolerante a nomes tipo: + base_CAM_A.bin + base_rgb.bin + base_cam0.bin + base_arquivo_CAM_A.bin + """ + if not bins_dir.is_dir(): + return None + + role_aliases = ROLE_ALIASES.get(role, [role]) + cam_key_variants = {cam_key, cam_key.upper(), cam_key.lower()} + + candidates = [] + for p in sorted(bins_dir.glob("*.bin")): + name = p.stem + name_l = name.lower() + base_l = base.lower() + if not name_l.startswith(base_l): continue - try: - count += augment_pair( - preview_path=preview_path, - mask_path=mask_path, - preview_out_dir=preview_out_dir, - mask_out_dir=mask_out_dir, - copies=copies, - ) + score = 0 + if name_l == base_l: + score += 1 + if any(v.lower() in name_l for v in cam_key_variants): + score += 10 + if any(a.lower() in name_l for a in role_aliases): + score += 12 - except Exception as e: - erros += 1 - print(f"[ERRO] [{group_name}] {base}: {e}") + # Compatibilidade com cam0/cam1/cam2 caso necessário. + if role == "rgb" and "cam0" in name_l: + score += 3 + if role == "re" and "cam1" in name_l: + score += 3 + if role == "nir" and "cam2" in name_l: + score += 3 - print( - f"[OK] Grupo '{group_name}' → {count} pares gerados. " - f"originais={len(bases)} | sem_mask={sem_mask} | erros={erros}" - ) + if score > 0: + candidates.append((score, p)) - return count + if not candidates: + return None + + candidates.sort(key=lambda x: x[0], reverse=True) + return candidates[0][1] -def main( - copies=5, - groups_csv=None, - src_root=DEFAULT_SRC_ROOT, - dst_root=DEFAULT_DST_ROOT, - limit=None, - seed=42, - clear_dst=False, -): - total = 0 +def discover_sample_cameras(meta: Dict[str, Any], bins_dir: Path, base: str) -> Dict[str, CameraBin]: + cam_info = extract_camera_info(meta) + out: Dict[str, CameraBin] = {} - print(f"[INFO] MODELO={MODELO}") - print(f"[INFO] SRC_ROOT={src_root}") - print(f"[INFO] DST_ROOT={dst_root}") - print(f"[INFO] copies={copies}") - print(f"[INFO] single_head=True") + # Primeiro caminho: usar camera_info do meta. + for cam_key, info in cam_info.items(): + role = normalize_role(info.get("role") or info.get("papel") or info.get("type"), cam_key) + if role not in ("rgb", "re", "nir"): + role = role_from_text(cam_key) or role_from_text(json.dumps(info, ensure_ascii=False)) - if not os.path.isdir(src_root): - raise SystemExit(f"[ERRO] src-root não encontrado: {src_root}") + if role not in ("rgb", "re", "nir"): + continue - if clear_dst: - print(f"[INFO] Limpando destino: {dst_root}") - maybe_clear_dir(dst_root) + width = get_int_any(info, ["width", "w", "sensor_width", "cols", "shape_w"]) + height = get_int_any(info, ["height", "h", "sensor_height", "rows", "shape_h"]) - grupos = list_groups(src_root) + if width is None or height is None: + # Fallback para campos globais do meta. + width = get_int_any(meta, ["width", "sensor_width", "raw_width", "w"]) + height = get_int_any(meta, ["height", "sensor_height", "raw_height", "h"]) - if groups_csv: - want = {g.strip() for g in groups_csv.split(",") if g.strip()} - grupos = [g for g in grupos if g in want] + if width is None or height is None: + continue - if not grupos: - print("[WARN] Nenhum grupo válido encontrado após filtro.") + bit_depth = get_int_any(info, ["bit_depth", "bits", "raw_bit_depth"], DEFAULT_BIT_DEPTH) + raw_format = get_str_any(info, ["raw_format", "format", "encoding"], DEFAULT_RAW_FORMAT) + bayer_pattern = get_str_any(info, ["bayer_pattern", "bayer", "pattern"], DEFAULT_BAYER_PATTERN) - if not grupos: - print("[WARN] Nenhum grupo encontrado.") - return + bin_path = find_bin_for_camera(bins_dir, base, role, cam_key) + if bin_path is None: + continue - print(f"Grupos encontrados: {', '.join(grupos)}") - - for g in grupos: - total += process_group( - src_root=src_root, - dst_root=dst_root, - group_name=g, - copies=copies, - limit=limit, - seed=seed, + out[role] = CameraBin( + role=role, + cam_key=str(cam_key), + path=bin_path, + width=int(width), + height=int(height), + bit_depth=int(bit_depth or DEFAULT_BIT_DEPTH), + raw_format=raw_format or DEFAULT_RAW_FORMAT, + bayer_pattern=bayer_pattern or DEFAULT_BAYER_PATTERN, ) - print(f"\nAugmentation completed! Total: {total} pares gerados.") + # Segundo caminho: se o meta não ajudou, tenta descobrir por nome. + if len(out) < 3: + global_w = get_int_any(meta, ["width", "sensor_width", "raw_width", "w"]) + global_h = get_int_any(meta, ["height", "sensor_height", "raw_height", "h"]) + + for role, cam_key in [("rgb", "CAM_A"), ("re", "CAM_B"), ("nir", "CAM_C")]: + if role in out: + continue + if global_w is None or global_h is None: + continue + bin_path = find_bin_for_camera(bins_dir, base, role, cam_key) + if bin_path is None: + continue + out[role] = CameraBin( + role=role, + cam_key=cam_key, + path=bin_path, + width=int(global_w), + height=int(global_h), + bit_depth=DEFAULT_BIT_DEPTH, + raw_format=DEFAULT_RAW_FORMAT, + bayer_pattern=DEFAULT_BAYER_PATTERN, + ) + + return out -if __name__ == "__main__": - ap = argparse.ArgumentParser( - description="Augmentação por grupos OAK-FCC-3 usando previews/masks, single head." +# ============================================================ +# Preview RGB simples a partir de RAW Bayer +# ============================================================ + +def _normalize_preview_channel(x: np.ndarray, lo_p: float = 1.0, hi_p: float = 99.5, gamma: float = 1.0 / 2.2) -> np.ndarray: + """ + Normalização visual robusta para preview, não científica. + Transforma um canal RAW/float em uint8 bonito para inspeção humana. + """ + y = x.astype(np.float32) + lo = float(np.percentile(y, lo_p)) + hi = float(np.percentile(y, hi_p)) + if hi <= lo: + hi = lo + 1.0 + y = np.clip((y - lo) / (hi - lo), 0.0, 1.0) + if gamma and abs(gamma - 1.0) > 1e-6: + y = np.power(y, gamma) + return np.clip(y * 255.0, 0, 255).astype(np.uint8) + + +def _get_bayer_cv2_code_like_core(bayer_pattern: str | None, algorithm: str = "ea") -> int: + """ + Replica o mapeamento usado pelo RawProcessorCore._get_bayer_cv2_code. + + Observação importante: + O mapeamento OpenCV parece invertido à primeira vista, mas é exatamente + o contrato usado no core atual para gerar o RGB do treinamento/inferência. + """ + p = str(bayer_pattern or DEFAULT_BAYER_PATTERN or "RGGB").upper() + algo = str(algorithm or "ea").lower() + + if algo in ("bilinear", "linear", "fast", "normal"): + code_map = { + "BGGR": cv2.COLOR_BayerRG2RGB, + "RGGB": cv2.COLOR_BayerBG2RGB, + "GRBG": cv2.COLOR_BayerGR2RGB, + "GBRG": cv2.COLOR_BayerGB2RGB, + } + elif algo in ("ea", "edge_aware", "edge-aware"): + code_map = { + "BGGR": cv2.COLOR_BayerRG2RGB_EA, + "RGGB": cv2.COLOR_BayerBG2RGB_EA, + "GRBG": cv2.COLOR_BayerGR2RGB_EA, + "GBRG": cv2.COLOR_BayerGB2RGB_EA, + } + else: + raise ValueError(f"demosaic_algorithm inválido: {algorithm}") + + if p not in code_map: + raise ValueError(f"Padrão Bayer não suportado para demosaic: {p}") + + return code_map[p] + + +def demosaic_raw10_to_rgb01(raw: np.ndarray, bayer_pattern: str = DEFAULT_BAYER_PATTERN, algorithm: str = "ea") -> np.ndarray: + """ + RAW Bayer 0..1023 -> RGB HWC float32 0..1 usando o mesmo mapeamento do core. + """ + raw_u16 = np.clip(raw, 0, 1023).astype(np.uint16) + cv_code = _get_bayer_cv2_code_like_core(bayer_pattern, algorithm=algorithm) + rgb16 = cv2.cvtColor(raw_u16, cv_code) + rgb = rgb16.astype(np.float32) / 1023.0 + return np.clip(rgb, 0.0, 1.0).astype(np.float32, copy=False) + + +def remosaic_rgb01_to_bayer_raw10(rgb: np.ndarray, bayer_pattern: str = DEFAULT_BAYER_PATTERN) -> np.ndarray: + """ + RGB HWC float32 0..1 -> mosaico Bayer RAW10 uint16. + + Isso é necessário porque NÃO podemos aplicar warp diretamente no mosaico Bayer. + O fluxo correto para augmentar RGB RAW é: + raw Bayer -> demosaic RGB -> augmentation -> remosaic Bayer -> pack RAW10. + + Não é uma reconstrução óptica perfeita, mas preserva o contrato RAW10 Bayer para + o normalize/RawProcessorCore e evita o xadrez colorido causado por interpolar + diretamente o mosaico. + """ + rgb = np.clip(rgb.astype(np.float32), 0.0, 1.0) + h, w = rgb.shape[:2] + raw = np.empty((h, w), dtype=np.float32) + + p = str(bayer_pattern or DEFAULT_BAYER_PATTERN or "RGGB").upper() + + if p == "RGGB": + raw[0::2, 0::2] = rgb[0::2, 0::2, 0] # R + raw[0::2, 1::2] = rgb[0::2, 1::2, 1] # G + raw[1::2, 0::2] = rgb[1::2, 0::2, 1] # G + raw[1::2, 1::2] = rgb[1::2, 1::2, 2] # B + elif p == "BGGR": + raw[0::2, 0::2] = rgb[0::2, 0::2, 2] # B + raw[0::2, 1::2] = rgb[0::2, 1::2, 1] # G + raw[1::2, 0::2] = rgb[1::2, 0::2, 1] # G + raw[1::2, 1::2] = rgb[1::2, 1::2, 0] # R + elif p == "GRBG": + raw[0::2, 0::2] = rgb[0::2, 0::2, 1] # G + raw[0::2, 1::2] = rgb[0::2, 1::2, 0] # R + raw[1::2, 0::2] = rgb[1::2, 0::2, 2] # B + raw[1::2, 1::2] = rgb[1::2, 1::2, 1] # G + elif p == "GBRG": + raw[0::2, 0::2] = rgb[0::2, 0::2, 1] # G + raw[0::2, 1::2] = rgb[0::2, 1::2, 2] # B + raw[1::2, 0::2] = rgb[1::2, 0::2, 0] # R + raw[1::2, 1::2] = rgb[1::2, 1::2, 1] # G + else: + raise ValueError(f"Padrão Bayer não suportado para remosaic: {p}") + + return np.clip(np.rint(raw * 1023.0), 0, 1023).astype(np.uint16) + + +def rgb01_to_preview_rgb8(rgb: np.ndarray) -> np.ndarray: + rgb = np.clip(rgb.astype(np.float32), 0.0, 1.0) + rgb8 = np.zeros(rgb.shape, dtype=np.uint8) + for c in range(3): + rgb8[:, :, c] = _normalize_preview_channel(rgb[:, :, c], lo_p=1.0, hi_p=99.5, gamma=1.0 / 2.2) + return rgb8 + + +def debayer_raw10_to_rgb8(raw: np.ndarray, bayer_pattern: str = DEFAULT_BAYER_PATTERN) -> np.ndarray: + """ + Preview visual a partir do RAW Bayer, usando o mesmo mapeamento Bayer do core. + """ + rgb01 = demosaic_raw10_to_rgb01(raw, bayer_pattern=bayer_pattern, algorithm="ea") + return rgb01_to_preview_rgb8(rgb01) + + +def _resolve_module_params_path(meta: Dict[str, Any]) -> Optional[str]: + candidates = [] + + for key in ("camera_params_json", "module_params_json"): + v = meta.get(key) + if v: + candidates.append(str(v)) + + stream_meta = meta.get("stream_meta") if isinstance(meta.get("stream_meta"), dict) else None + if stream_meta: + for key in ("camera_params_json", "module_params_json"): + v = stream_meta.get(key) + if v: + candidates.append(str(v)) + + cfg_module = read_config_value("module_params_json", None) + if cfg_module: + candidates.append(str(cfg_module)) + + candidates.append("calibration/module_params.json") + + for c in candidates: + p = Path(c) + if p.is_file(): + return str(p) + + return None + + +def _build_processing_meta_for_core(meta: Dict[str, Any], cameras: Dict[str, CameraBin]) -> Dict[str, Any]: + """ + Monta o meta que o RawProcessorCore espera no build_infer_tensor_from_stream. + O validador faz isso a partir de stream_meta + camera_info. Aqui criamos/atualizamos + esse pacote usando os bins augmentados. + """ + stream_meta = copy.deepcopy(meta.get("stream_meta") if isinstance(meta.get("stream_meta"), dict) else {}) + + camera_info = stream_meta.get("camera_info") + if not isinstance(camera_info, dict): + camera_info = copy.deepcopy(meta.get("camera_info") if isinstance(meta.get("camera_info"), dict) else {}) + + camera_frames = stream_meta.get("camera_frames") + if not isinstance(camera_frames, dict): + camera_frames = copy.deepcopy(meta.get("camera_frames") if isinstance(meta.get("camera_frames"), dict) else {}) + + for role, cam in cameras.items(): + info = dict(camera_info.get(cam.cam_key, {}) or {}) + info.update({ + "role": role, + "width": int(cam.width), + "height": int(cam.height), + "bit_depth": int(cam.bit_depth), + "raw_format": cam.raw_format or DEFAULT_RAW_FORMAT, + "packed": True, + "channels": 1, + "bayer_pattern": cam.bayer_pattern or DEFAULT_BAYER_PATTERN, + }) + camera_info[cam.cam_key] = info + camera_frames[cam.cam_key] = dict(info) + + stream_meta["frame_type"] = "RAW_BRUTO" + stream_meta["camera_info"] = camera_info + stream_meta["camera_frames"] = camera_frames + stream_meta["payload_sources"] = [cameras[r].cam_key for r in ("rgb", "re", "nir") if r in cameras] + + # A normalização radiométrica do core busca frame_controls no próprio meta. + # Preservamos os controles originais se existirem. + for key in ("frame_controls", "actual_camera_controls", "startup_camera_controls"): + if key not in stream_meta and isinstance(meta.get(key), dict): + stream_meta[key] = copy.deepcopy(meta[key]) + + return stream_meta + + +def get_preview_core_cached(sensor_width: int, sensor_height: int, bayer: str, calib_path: str): + """ + Reutiliza o RawProcessorCore entre amostras para evitar recarregar flat-field + e recriar caches de gain/remap toda hora. + """ + if not _HAS_RAW_PROCESSOR_CORE: + return None + + key = ( + int(sensor_width), + int(sensor_height), + str(bayer).upper(), + str(Path(calib_path).resolve()) if calib_path else "", ) - ap.add_argument("--copies", type=int, default=5, help="Número de cópias augmentadas por imagem.") + core = _CORE_PREVIEW_CACHE.get(key) + if core is not None: + return core + + core = RawProcessorCore( + sensor_width=int(sensor_width), + sensor_height=int(sensor_height), + bayer_pattern=str(bayer).upper(), + calibration_json_path=calib_path, + ) + + _CORE_PREVIEW_CACHE[key] = core + return core + + +def build_core_preview_from_augmented_raws(meta: Dict[str, Any], cameras: Dict[str, CameraBin]) -> Optional[np.ndarray]: + """ + Gera o preview salvo passando os bins augmentados pelo RawProcessorCore, + igual ao check_saved_files faz para montar o MULTISPEC final. + + Saída: RGB uint8 do tensor final [R,G,B], pronto para salvar com PIL. + """ + if not _HAS_RAW_PROCESSOR_CORE: + return None + + required = ("rgb", "re", "nir") + if any(r not in cameras or cameras[r].data is None for r in required): + return None + + sensor_width = int(meta.get("sensor_width") or cameras["rgb"].width) + sensor_height = int(meta.get("sensor_height") or cameras["rgb"].height) + bayer = str(meta.get("bayer_pattern") or cameras["rgb"].bayer_pattern or DEFAULT_BAYER_PATTERN) + calib_path = _resolve_module_params_path(meta) + + if not calib_path: + return None + + frame = {} + for role in required: + cam = cameras[role] + frame[cam.cam_key] = pack_raw10_packed_array(cam.data) + + processing_meta = _build_processing_meta_for_core(meta, cameras) + + core = get_preview_core_cached( + sensor_width=sensor_width, + sensor_height=sensor_height, + bayer=bayer, + calib_path=calib_path, + ) + + if core is None: + return None + + tensor = core.build_infer_tensor_from_stream(frame, processing_meta, 5) + if tensor is None or not isinstance(tensor, np.ndarray) or tensor.ndim != 3 or tensor.shape[0] < 3: + return None + + rgb_hwc = np.transpose(tensor[:3].astype(np.float32), (1, 2, 0)) + rgb8 = np.clip(rgb_hwc * 255.0, 0, 255).astype(np.uint8) + return rgb8 + + +def build_beauty_preview_from_augmented_raws(cameras: Dict[str, CameraBin]) -> np.ndarray: + """ + Fallback visual a partir dos RAWs augmentados, sem passar pelo core. + Preferimos build_core_preview_from_augmented_raws sempre que possível. + """ + rgb_cam = cameras.get("rgb") + if rgb_cam is None or rgb_cam.data is None: + raise RuntimeError("Sem RAW RGB para gerar preview.") + + return debayer_raw10_to_rgb8(rgb_cam.data, rgb_cam.bayer_pattern) + + +def save_preview_from_augmented_raws(path: Path, cameras: Dict[str, CameraBin], meta: Optional[Dict[str, Any]] = None) -> str: + ensure_dir(path.parent) + + method = "fallback_cam_a_debayer_preview" + rgb8 = None + + if meta is not None: + try: + rgb8 = build_core_preview_from_augmented_raws(meta, cameras) + if rgb8 is not None: + method = "raw_processor_core_multispec_rgb_final" + except Exception as e: + print(f"[WARN] Preview via RawProcessorCore falhou em {path.name}: {e}. Usando fallback CAM_A.") + rgb8 = None + + if rgb8 is None: + rgb8 = build_beauty_preview_from_augmented_raws(cameras) + + Image.fromarray(rgb8).save(path) + return method + + +# ============================================================ +# Augmentation +# ============================================================ + +def sample_aug_params(rng: random.Random, seed_value: int) -> AugParams: + # Geometria conservadora. + perspective = rng.random() < 0.15 + blur_enabled = rng.random() < 0.15 + + return AugParams( + seed=seed_value, + flip_h=rng.random() < 0.50, + shift_x_frac=rng.uniform(-0.015, 0.015), + shift_y_frac=rng.uniform(-0.015, 0.015), + scale=rng.uniform(0.92, 1.08), + rotate_deg=rng.uniform(-4.0, 4.0), + perspective=perspective, + perspective_strength=rng.uniform(0.002, 0.012) if perspective else 0.0, + exposure_mult=rng.uniform(0.75, 1.30), + gamma=rng.uniform(0.92, 1.08), + rgb_channel_gain=( + rng.uniform(0.92, 1.08), + rng.uniform(0.92, 1.08), + rng.uniform(0.92, 1.08), + ), + re_gain=rng.uniform(0.85, 1.20), + nir_gain=rng.uniform(0.85, 1.20), + shadow_enabled=rng.random() < 0.25, + shadow_strength=rng.uniform(0.12, 0.35), + shadow_angle_deg=rng.uniform(0.0, 180.0), + highlight_enabled=rng.random() < 0.15, + highlight_strength=rng.uniform(0.05, 0.18), + noise_enabled=rng.random() < 0.30, + noise_sigma_dn=rng.uniform(1.5, 6.0), + blur_enabled=blur_enabled, + blur_kernel=3 if blur_enabled else 1, + ) + + +def build_affine_matrix(width: int, height: int, p: AugParams) -> np.ndarray: + cx = width * 0.5 + cy = height * 0.5 + M = cv2.getRotationMatrix2D((cx, cy), p.rotate_deg, p.scale) + M[0, 2] += p.shift_x_frac * width + M[1, 2] += p.shift_y_frac * height + return M + + +def warp_image(img: np.ndarray, M: np.ndarray, out_size: Tuple[int, int], is_mask: bool) -> np.ndarray: + interp = cv2.INTER_NEAREST if is_mask else cv2.INTER_LINEAR + return cv2.warpAffine( + img, + M, + out_size, + flags=interp, + borderMode=cv2.BORDER_REFLECT_101, + ) + + +def build_perspective_matrix(width: int, height: int, p: AugParams, rng: random.Random) -> Optional[np.ndarray]: + if not p.perspective: + return None + + s = p.perspective_strength + dx = width * s + dy = height * s + + src = np.float32([ + [0, 0], + [width - 1, 0], + [width - 1, height - 1], + [0, height - 1], + ]) + + dst = src + np.float32([ + [rng.uniform(-dx, dx), rng.uniform(-dy, dy)], + [rng.uniform(-dx, dx), rng.uniform(-dy, dy)], + [rng.uniform(-dx, dx), rng.uniform(-dy, dy)], + [rng.uniform(-dx, dx), rng.uniform(-dy, dy)], + ]) + + return cv2.getPerspectiveTransform(src, dst) + + +def warp_perspective(img: np.ndarray, H: np.ndarray, out_size: Tuple[int, int], is_mask: bool) -> np.ndarray: + interp = cv2.INTER_NEAREST if is_mask else cv2.INTER_LINEAR + return cv2.warpPerspective( + img, + H, + out_size, + flags=interp, + borderMode=cv2.BORDER_REFLECT_101, + ) + + +def apply_geometry_raw(raw: np.ndarray, p: AugParams, rng: random.Random) -> np.ndarray: + h, w = raw.shape[:2] + + out = raw + if p.flip_h: + out = cv2.flip(out, 1) + + M = build_affine_matrix(w, h, p) + out = warp_image(out, M, (w, h), is_mask=False) + + H = build_perspective_matrix(w, h, p, rng) + if H is not None: + out = warp_perspective(out, H, (w, h), is_mask=False) + + return out + + +def apply_geometry_mask(mask: np.ndarray, p: AugParams, rng: random.Random) -> np.ndarray: + h, w = mask.shape[:2] + + out = mask + if p.flip_h: + out = cv2.flip(out, 1) + + M = build_affine_matrix(w, h, p) + out = warp_image(out, M, (w, h), is_mask=True) + + H = build_perspective_matrix(w, h, p, rng) + if H is not None: + out = warp_perspective(out, H, (w, h), is_mask=True) + + return out + + +def make_spatial_light_map(shape: Tuple[int, int], p: AugParams) -> np.ndarray: + h, w = shape + yy, xx = np.mgrid[0:h, 0:w].astype(np.float32) + xx = (xx / max(w - 1, 1)) - 0.5 + yy = (yy / max(h - 1, 1)) - 0.5 + + light = np.ones((h, w), dtype=np.float32) + + if p.shadow_enabled: + theta = math.radians(p.shadow_angle_deg) + direction = math.cos(theta) * xx + math.sin(theta) * yy + direction = (direction - direction.min()) / max(direction.max() - direction.min(), 1e-6) + shadow = 1.0 - p.shadow_strength * direction + light *= shadow + + if p.highlight_enabled: + # Mancha larga e suave, simulando região de sol/reflexo. + cx = np.random.uniform(-0.25, 0.25) + cy = np.random.uniform(-0.25, 0.25) + sigma = np.random.uniform(0.20, 0.42) + d2 = (xx - cx) ** 2 + (yy - cy) ** 2 + blob = np.exp(-d2 / (2 * sigma * sigma)) + light *= 1.0 + p.highlight_strength * blob + + return np.clip(light, 0.50, 1.45).astype(np.float32) + + +def apply_gamma_raw(raw: np.ndarray, gamma: float) -> np.ndarray: + if abs(gamma - 1.0) < 1e-3: + return raw + x = np.clip(raw.astype(np.float32) / 1023.0, 0, 1) + x = np.power(x, gamma) + return x * 1023.0 + + +def apply_bayer_channel_gains(raw: np.ndarray, gains_rgb: Tuple[float, float, float], pattern: str) -> np.ndarray: + """ + Aplica ganhos aproximados por canal em mosaico Bayer. + Para preview/treino, isso simula variação de balanço/cor no sensor RGB. + """ + out = raw.astype(np.float32).copy() + r_gain, g_gain, b_gain = gains_rgb + pat = (pattern or DEFAULT_BAYER_PATTERN).upper() + + if pat == "RGGB": + out[0::2, 0::2] *= r_gain + out[0::2, 1::2] *= g_gain + out[1::2, 0::2] *= g_gain + out[1::2, 1::2] *= b_gain + elif pat == "BGGR": + out[0::2, 0::2] *= b_gain + out[0::2, 1::2] *= g_gain + out[1::2, 0::2] *= g_gain + out[1::2, 1::2] *= r_gain + elif pat == "GRBG": + out[0::2, 0::2] *= g_gain + out[0::2, 1::2] *= r_gain + out[1::2, 0::2] *= b_gain + out[1::2, 1::2] *= g_gain + elif pat == "GBRG": + out[0::2, 0::2] *= g_gain + out[0::2, 1::2] *= b_gain + out[1::2, 0::2] *= r_gain + out[1::2, 1::2] *= g_gain + else: + out *= float(np.mean(gains_rgb)) + + return out + + +def apply_radiometry_rgb01(rgb: np.ndarray, p: AugParams, rng_np: np.random.Generator, light_map_cache: Dict[Tuple[int, int], np.ndarray]) -> np.ndarray: + """ + Radiometria para RGB já demosaicado. + Evita mexer no mosaico Bayer diretamente. + """ + h, w = rgb.shape[:2] + x = np.clip(rgb.astype(np.float32), 0.0, 1.0) + + x *= np.float32(p.exposure_mult) + + gains = np.array(p.rgb_channel_gain, dtype=np.float32).reshape(1, 1, 3) + x *= gains + + key = (h, w) + if key not in light_map_cache: + light_map_cache[key] = make_spatial_light_map((h, w), p) + x *= light_map_cache[key][:, :, None] + + if abs(p.gamma - 1.0) > 1e-3: + x = np.power(np.clip(x, 0.0, 1.0), p.gamma) + + if p.blur_enabled and p.blur_kernel >= 3: + x = cv2.GaussianBlur(x, (p.blur_kernel, p.blur_kernel), 0) + + if p.noise_enabled: + # Converte sigma DN RAW10 para escala 0..1. + sigma01 = float(p.noise_sigma_dn) / 1023.0 + noise = rng_np.normal(0.0, sigma01, size=x.shape).astype(np.float32) + x += noise + + return np.clip(x, 0.0, 1.0).astype(np.float32, copy=False) + + +def apply_radiometry_raw(raw: np.ndarray, role: str, p: AugParams, cam: CameraBin, rng_np: np.random.Generator, light_map_cache: Dict[Tuple[int, int], np.ndarray]) -> np.ndarray: + h, w = raw.shape[:2] + x = raw.astype(np.float32) + + # Base global de exposição: muda todo mundo junto. + x *= p.exposure_mult + + # Variação por papel espectral/canal. + # OBS: RGB não deve passar por aqui, porque RGB Bayer precisa ser demosaicado + # antes de qualquer warp/radiometria por canal. + if role == "re": + x *= p.re_gain + elif role == "nir": + x *= p.nir_gain + + # Mesmo mapa espacial de luz por tamanho, para manter coerência física. + key = (h, w) + if key not in light_map_cache: + light_map_cache[key] = make_spatial_light_map((h, w), p) + x *= light_map_cache[key] + + # Gamma leve, opcional. Em RAW científico puro, gamma seria discutível. + # Mantemos bem pequeno para simular resposta/exposição não ideal. + x = apply_gamma_raw(x, p.gamma) + + if p.blur_enabled and p.blur_kernel >= 3: + x = cv2.GaussianBlur(x, (p.blur_kernel, p.blur_kernel), 0) + + if p.noise_enabled: + noise = rng_np.normal(0.0, p.noise_sigma_dn, size=x.shape).astype(np.float32) + x += noise + + return np.clip(x, 0, 1023).astype(np.uint16) + + +# ============================================================ +# Processamento de amostra/grupo +# ============================================================ + +def validate_mask_colors(mask_before: np.ndarray, mask_after: np.ndarray, name: str) -> None: + if mask_before.ndim == 2 or mask_after.ndim == 2: + before_colors = int(np.unique(mask_before.reshape(-1)).size) + after_colors = int(np.unique(mask_after.reshape(-1)).size) + else: + before_colors = len(set(map(tuple, mask_before.reshape(-1, mask_before.shape[2])))) + after_colors = len(set(map(tuple, mask_after.reshape(-1, mask_after.shape[2])))) + + if after_colors > max(before_colors * 3, 64): + print( + f"[WARN] {name}: máscara ganhou muitas cores. " + f"antes={before_colors}, depois={after_colors}. Confira interpolação NEAREST." + ) + + +def save_aug_meta( + meta: Dict[str, Any], + out_path: Path, + base: str, + new_base: str, + params: AugParams, + cameras: Dict[str, CameraBin], + preview_method: Optional[str] = None, +) -> None: + out = copy.deepcopy(meta) + out["synthetic"] = True + out["augmentation"] = { + "script": "_5_augmentation_raw_oak.py", + "parent_base": base, + "new_base": new_base, + "params": asdict(params), + } + + # Deixa o JSON compatível com o check_saved_files / validador. + out["saved_payload_type"] = "raw_native_multi" + out["saved_preview_method"] = preview_method or "unknown" + out["saved_payload_paths"] = {} + out["saved_payload_shapes"] = {} + out["saved_payload_dtypes"] = {} + + stream_meta = copy.deepcopy(out.get("stream_meta") if isinstance(out.get("stream_meta"), dict) else {}) + camera_info = stream_meta.get("camera_info") if isinstance(stream_meta.get("camera_info"), dict) else {} + camera_frames = stream_meta.get("camera_frames") if isinstance(stream_meta.get("camera_frames"), dict) else {} + + out["augmented_bins"] = {} + + for role, cam in cameras.items(): + packed_w = int(math.ceil(int(cam.width) * 10 / 8)) + fname = f"{new_base}_{cam.cam_key}.bin" + + out["saved_payload_paths"][cam.cam_key] = fname + out["saved_payload_shapes"][cam.cam_key] = [int(cam.height), int(packed_w)] + out["saved_payload_dtypes"][cam.cam_key] = "uint8" + + cam_info = dict(camera_info.get(cam.cam_key, {}) or {}) + cam_info.update({ + "role": role, + "width": int(cam.width), + "height": int(cam.height), + "bit_depth": int(cam.bit_depth), + "raw_format": cam.raw_format or DEFAULT_RAW_FORMAT, + "packed": True, + "channels": 1, + "bayer_pattern": cam.bayer_pattern or DEFAULT_BAYER_PATTERN, + }) + camera_info[cam.cam_key] = cam_info + camera_frames[cam.cam_key] = dict(cam_info) + + out["augmented_bins"][role] = { + "cam_key": cam.cam_key, + "filename": fname, + "role": role, + "width": int(cam.width), + "height": int(cam.height), + "bit_depth": int(cam.bit_depth), + "raw_format": cam.raw_format, + "bayer_pattern": cam.bayer_pattern, + } + + stream_meta["frame_type"] = "RAW_BRUTO" + stream_meta["camera_info"] = camera_info + stream_meta["camera_frames"] = camera_frames + stream_meta["payload_sources"] = [cameras[r].cam_key for r in ("rgb", "re", "nir") if r in cameras] + out["stream_meta"] = stream_meta + + save_json(out_path, out) + + +def process_one_sample( + base: str, + group_name: str, + src_group: Path, + dst_group: Path, + mask_path: Path, + meta_path: Path, + copies: int, + rng_global: random.Random, + dry_run: bool = False, +) -> Tuple[int, int]: + bins_dir = src_group / "bins" + meta = load_json(meta_path) + cameras = discover_sample_cameras(meta, bins_dir, base) + + required = ["rgb", "re", "nir"] + missing = [r for r in required if r not in cameras] + if missing: + print(f"[WARN] [{group_name}] {base}: câmeras ausentes {missing}. Pulando.") + return 0, 1 + + mask = load_mask(mask_path) + + # Carrega todos os RAWs uma vez. + for role, cam in cameras.items(): + cam.data = read_raw_bin(cam.path, cam.width, cam.height, cam.raw_format) + + generated = 0 + errors = 0 + + out_bins = dst_group / "bins" + out_masks = dst_group / "masks" + out_metas = dst_group / "metas" + out_previews = dst_group / "previews" + + for i in range(copies): + seed_value = rng_global.randint(0, 2**31 - 1) + rng = random.Random(seed_value) + rng_np = np.random.default_rng(seed_value) + params = sample_aug_params(rng, seed_value) + new_base = f"{base}_aug_{i:02d}" + + try: + if dry_run: + generated += 1 + continue + + # Mesma geometria para mask e raws. + mask_aug = apply_geometry_mask(mask, params, random.Random(seed_value + 1000)) + validate_mask_colors(mask, mask_aug, f"{new_base}_mask") + + light_map_cache: Dict[Tuple[int, int], np.ndarray] = {} + cams_out: Dict[str, CameraBin] = {} + + for role in required: + cam = cameras[role] + assert cam.data is not None + + if role == "rgb": + # RGB RAW é mosaico Bayer. Não podemos aplicar warp direto nele, + # porque isso mistura pixels R/G/B antes do demosaic e cria ruído colorido. + # Fluxo correto offline: + # RAW Bayer -> RGB linear -> geometria/radiometria -> mosaico Bayer -> RAW10 + rgb01 = demosaic_raw10_to_rgb01( + cam.data, + bayer_pattern=cam.bayer_pattern, + algorithm="ea", + ) + + rgb_geom = rgb01 + if params.flip_h: + rgb_geom = cv2.flip(rgb_geom, 1) + + h_rgb, w_rgb = rgb_geom.shape[:2] + M_rgb = build_affine_matrix(w_rgb, h_rgb, params) + rgb_geom = warp_image(rgb_geom, M_rgb, (w_rgb, h_rgb), is_mask=False) + + H_rgb = build_perspective_matrix(w_rgb, h_rgb, params, random.Random(seed_value + 1000)) + if H_rgb is not None: + rgb_geom = warp_perspective(rgb_geom, H_rgb, (w_rgb, h_rgb), is_mask=False) + + rgb_aug = apply_radiometry_rgb01( + rgb_geom, + params, + rng_np, + light_map_cache, + ) + + raw_aug = remosaic_rgb01_to_bayer_raw10( + rgb_aug, + bayer_pattern=cam.bayer_pattern, + ) + + else: + raw_geom = apply_geometry_raw(cam.data, params, random.Random(seed_value + 1000)) + raw_aug = apply_radiometry_raw(raw_geom, role, params, cam, rng_np, light_map_cache) + + out_cam = copy.deepcopy(cam) + out_cam.path = out_bins / f"{new_base}_{cam.cam_key}.bin" + out_cam.data = raw_aug + cams_out[role] = out_cam + + write_raw_bin(out_cam.path, raw_aug, cam.raw_format) + + # Máscara com mesmo nome-base. + mask_ext = mask_path.suffix.lower() if mask_path.suffix.lower() in IMG_EXTS else ".png" + save_mask(out_masks / f"{new_base}{mask_ext}", mask_aug) + + # Preview visual usando o mesmo caminho do validador/check_saved_files: + # bins augmentados -> RawProcessorCore -> MULTISPEC final -> RGB final. + # Se o core falhar, cai no fallback CAM_A e avisa no console. + preview_method = save_preview_from_augmented_raws( + out_previews / f"{new_base}.png", + cams_out, + meta=meta, + ) + + # Meta novo, com mesmo stem do preview/mask para manter o layout dataset limpo. + save_aug_meta( + meta=meta, + out_path=out_metas / f"{new_base}.json", + base=base, + new_base=new_base, + params=params, + cameras=cams_out, + preview_method=preview_method, + ) + + generated += 1 + + except Exception as e: + errors += 1 + print(f"[ERRO] [{group_name}] {new_base}: {e}") + + return generated, errors + + +def parse_group_copies(text: Optional[str]) -> Dict[str, int]: + """ + Ex: chao:1,chao_cana:3,chao_erva:3,chao_cana_erva:5 + """ + out: Dict[str, int] = {} + if not text: + return out + + for part in text.split(","): + part = part.strip() + if not part: + continue + if ":" not in part: + raise ValueError(f"group-copies inválido: {part}. Use grupo:N") + g, n = part.split(":", 1) + out[g.strip()] = int(n.strip()) + return out + + +def process_group( + src_root: Path, + dst_root: Path, + group_name: str, + copies: int, + limit: Optional[int], + seed: int, + dry_run: bool, +) -> Tuple[int, int, int]: + src_group = src_root / group_name + dst_group = dst_root / group_name + + mask_map = map_files_by_base(src_group / "masks", IMG_EXTS) + meta_map = map_files_by_base(src_group / "metas", META_EXTS) + + bases = sorted(set(mask_map.keys()) & set(meta_map.keys())) + + if limit is not None and limit > 0 and limit < len(bases): + rng_select = random.Random(seed) + bases = sorted(rng_select.sample(bases, limit)) + print(f"[INFO] [{group_name}] limit={limit}, amostras selecionadas={len(bases)}") + + if not bases: + print(f"[WARN] [{group_name}] Nenhum par mask/meta encontrado.") + return 0, 0, 0 + + ensure_dir(dst_group / "bins") + ensure_dir(dst_group / "masks") + ensure_dir(dst_group / "metas") + ensure_dir(dst_group / "previews") + + rng_global = random.Random(seed + abs(hash(group_name)) % 1000000) + + total_generated = 0 + total_errors = 0 + + for idx, base in enumerate(bases, start=1): + gen, err = process_one_sample( + base=base, + group_name=group_name, + src_group=src_group, + dst_group=dst_group, + mask_path=mask_map[base], + meta_path=meta_map[base], + copies=copies, + rng_global=rng_global, + dry_run=dry_run, + ) + total_generated += gen + total_errors += err + + if idx % 25 == 0 or idx == len(bases): + print( + f"[INFO] [{group_name}] {idx}/{len(bases)} amostras | " + f"gerados={total_generated} | erros={total_errors}" + ) + + print( + f"[OK] Grupo '{group_name}' concluído: " + f"originais={len(bases)} | copies={copies} | gerados={total_generated} | erros={total_errors}" + ) + + return len(bases), total_generated, total_errors + + +# ============================================================ +# CLI +# ============================================================ + +def main() -> None: + dataset_base = Path("dataset") + + default_src = dataset_base / "original" / "group" + if not default_src.exists(): + # Compatibilidade com script antigo que usava "originals". + alt = dataset_base / "originals" / "group" + if alt.exists(): + default_src = alt + + default_dst = dataset_base / "augmented" / "group" + + ap = argparse.ArgumentParser( + description="Augmentation multiespectral RAW10 packed para dataset OAK-FCC-3." + ) + + ap.add_argument("--copies", type=int, default=3, help="Cópias augmentadas por amostra, se --group-copies não sobrescrever.") + ap.add_argument("--group-copies", type=str, default=None, help="Cópias por grupo. Ex: chao:1,chao_cana:3,chao_erva:3,chao_cana_erva:5") ap.add_argument("--groups", type=str, default=None, help="Lista de grupos separados por vírgula.") - ap.add_argument("--src-root", type=str, default=DEFAULT_SRC_ROOT, help="Raiz dos grupos originais.") - ap.add_argument("--dst-root", type=str, default=DEFAULT_DST_ROOT, help="Raiz dos grupos augmentados.") - ap.add_argument("--limit", type=int, default=None, help="Quantidade máxima de amostras originais por grupo para augmentar.") - ap.add_argument("--seed", type=int, default=42, help="Seed para seleção reproduzível quando usar --limit.") + ap.add_argument("--src-root", type=str, default=str(default_src), help="Raiz dos grupos originais.") + ap.add_argument("--dst-root", type=str, default=str(default_dst), help="Raiz dos grupos augmentados.") + ap.add_argument("--limit", type=int, default=None, help="Quantidade máxima de amostras originais por grupo.") + ap.add_argument("--seed", type=int, default=42, help="Seed geral para reproducibilidade.") ap.add_argument("--clear-dst", action="store_true", help="Apaga dst-root antes de gerar.") + ap.add_argument("--dry-run", action="store_true", help="Simula sem salvar arquivos.") args = ap.parse_args() - random.seed(args.seed) + src_root = Path(args.src_root) + dst_root = Path(args.dst_root) + group_copies = parse_group_copies(args.group_copies) - main( - copies=args.copies, - groups_csv=args.groups, - src_root=args.src_root, - dst_root=args.dst_root, - limit=args.limit, - seed=args.seed, - clear_dst=args.clear_dst, - ) \ No newline at end of file + print("==============================================") + print("Augmentation RAW OAK-FCC-3") + print(f"MODEL/DATASET : {dataset_base}") + print(f"SRC_ROOT : {src_root}") + print(f"DST_ROOT : {dst_root}") + print(f"copies : {args.copies}") + print(f"group_copies : {group_copies if group_copies else '{}'}") + print(f"limit : {args.limit}") + print(f"seed : {args.seed}") + print(f"dry_run : {args.dry_run}") + print("==============================================") + + if not src_root.exists(): + raise SystemExit(f"[ERRO] src-root não encontrado: {src_root}") + + if args.clear_dst and not args.dry_run: + print(f"[INFO] Limpando destino: {dst_root}") + clear_dir(dst_root) + else: + ensure_dir(dst_root) + + groups = list_groups(src_root) + + if args.groups: + wanted = {g.strip() for g in args.groups.split(",") if g.strip()} + groups = [g for g in groups if g in wanted] + + if not groups: + raise SystemExit("[WARN] Nenhum grupo válido encontrado.") + + print(f"Grupos encontrados: {', '.join(groups)}") + + total_originals = 0 + total_generated = 0 + total_errors = 0 + + for group_name in groups: + copies = group_copies.get(group_name, args.copies) + if copies <= 0: + print(f"[INFO] [{group_name}] copies={copies}. Pulando.") + continue + + n_orig, n_gen, n_err = process_group( + src_root=src_root, + dst_root=dst_root, + group_name=group_name, + copies=copies, + limit=args.limit, + seed=args.seed, + dry_run=args.dry_run, + ) + total_originals += n_orig + total_generated += n_gen + total_errors += n_err + + print("\n==============================================") + print("Resumo final") + print(f"Originais processados : {total_originals}") + print(f"Amostras geradas : {total_generated}") + print(f"Erros/Pulos : {total_errors}") + print(f"Destino : {dst_root}") + print("==============================================") + + +if __name__ == "__main__": + main() diff --git a/Python/OAK/datasets/oak-fcc-3/_6_normalize.py b/Python/OAK/datasets/oak-fcc-3/_6_normalize.py index 840aae7e4..4d4246123 100644 --- a/Python/OAK/datasets/oak-fcc-3/_6_normalize.py +++ b/Python/OAK/datasets/oak-fcc-3/_6_normalize.py @@ -814,7 +814,11 @@ def main(): description="Normaliza RAW_BRUTO OAK-FCC-3 para tensor MULTISPEC final de treino." ) - ap.add_argument("--src-root", default="dataset/original/group") + ap.add_argument( + "--src-roots", + default="dataset/original/group;dataset/augmented/group", + help="Raízes de entrada separadas por ';'. Ex: dataset/original/group;dataset/augmented/group" + ) ap.add_argument("--out-root", default="dataset") ap.add_argument("--module-params", default=DEFAULT_MODULE_PARAMS) ap.add_argument("--res", default=f"{DEFAULT_RES[0]}x{DEFAULT_RES[1]}", help="Resolução final WxH.") @@ -827,11 +831,23 @@ def main(): args = ap.parse_args() - src_root = Path(args.src_root) + src_roots = [ + Path(x.strip()) + for x in str(args.src_roots).split(";") + if x.strip() + ] + out_dataset_root = Path(args.out_root) - if not src_root.is_dir(): - raise SystemExit(f"[ERRO] src-root não encontrado: {src_root}") + valid_src_roots = [] + for src_root in src_roots: + if src_root.is_dir(): + valid_src_roots.append(src_root) + else: + print(f"[WARN] src-root não encontrado, ignorando: {src_root}") + + if not valid_src_roots: + raise SystemExit(f"[ERRO] Nenhum src-root válido encontrado: {src_roots}") try: res_w, res_h = [int(x) for x in args.res.lower().split("x")] @@ -856,44 +872,53 @@ def main(): cor_para_id, _, _, ignore_rgb = carregar_labelmap_completo(str(labelmap_path)) ignore_id = _infer_ignore_id(ignore_rgb, 255) - all_groups = list_groups(src_root) + groups_by_root = [] - if args.groups: - want = {g.strip() for g in args.groups.split(",") if g.strip()} - all_groups = [g for g in all_groups if g in want] + for src_root in valid_src_roots: + groups = list_groups(src_root) - if not all_groups: + if args.groups: + want = {g.strip() for g in args.groups.split(",") if g.strip()} + groups = [g for g in groups if g in want] + + if groups: + groups_by_root.append((src_root, groups)) + + if not groups_by_root: raise SystemExit("[ERRO] Nenhum grupo encontrado.") print("============================================") print("Normalize OAK-FCC-3") - print(f"SRC : {src_root}") + print(f"SRC : {[str(x) for x in valid_src_roots]}") print(f"OUT : {output_root}") print(f"MODULE PARAM : {args.module_params}") print(f"RES : {res}") print(f"RAW SIZE : {raw_size}") - print(f"GROUPS : {', '.join(all_groups)}") + print("GROUPS :") + for root, groups in groups_by_root: + print(f" - {root}: {', '.join(groups)}") print(f"SKIP BAD : {args.skip_bad_quality}") print("============================================") running_stats = RunningStats() all_rows = [] - for group_name in all_groups: - rows = process_group( - group_name=group_name, - src_root=src_root, - output_root=output_root, - res=res, - raw_size=raw_size, - module_params_arg=args.module_params, - cor_para_id=cor_para_id, - ignore_id=ignore_id, - stats=running_stats, - dataset_root=out_dataset_root, - skip_bad_quality=args.skip_bad_quality, - ) - all_rows.extend(rows) + for src_root, all_groups in groups_by_root: + for group_name in all_groups: + rows = process_group( + group_name=group_name, + src_root=src_root, + output_root=output_root, + res=res, + raw_size=raw_size, + module_params_arg=args.module_params, + cor_para_id=cor_para_id, + ignore_id=ignore_id, + stats=running_stats, + dataset_root=out_dataset_root, + skip_bad_quality=args.skip_bad_quality, + ) + all_rows.extend(rows) manifest_path = Path(args.manifest) if args.manifest else output_root / "normalize_manifest.csv" write_manifest(manifest_path, all_rows) diff --git a/Python/OAK/datasets/oak-fcc-3/_9_test_multihead.py b/Python/OAK/datasets/oak-fcc-3/_9_test_multihead.py index ac337a511..c0f51d0f8 100644 --- a/Python/OAK/datasets/oak-fcc-3/_9_test_multihead.py +++ b/Python/OAK/datasets/oak-fcc-3/_9_test_multihead.py @@ -1043,6 +1043,241 @@ class MultiHeadTester: return preds, probs, t_ms +class OnnxMultiHeadTester: + def __init__( + self, + config: dict, + onnx_path: Path, + provider: str, + channels: int, + heads_config: Dict[str, dict], + mean: Optional[Sequence[float]], + std: Optional[Sequence[float]], + trt_home: Optional[str] = None, + trt_fp16: bool = True, + ): + self.config = config + self.onnx_path = onnx_path + self.provider = provider.lower() + self.channels = int(channels) + self.heads_config = heads_config + self.runtime_mode = str(config.get("runtime_mode", "all")).lower() + + self.mean = None if mean is None else np.asarray(mean, dtype=np.float32).reshape(1, channels, 1, 1) + self.std = None if std is None else np.asarray(std, dtype=np.float32).reshape(1, channels, 1, 1) + + print(f"[ONNX_MODEL] onnx={onnx_path}") + print(f"[ONNX_MODEL] provider={provider}") + + self.session = self._create_session( + onnx_path=onnx_path, + provider=provider, + trt_home=trt_home, + trt_fp16=trt_fp16, + ) + + self.input_name = self.session.get_inputs()[0].name + self.output_names = [o.name for o in self.session.get_outputs()] + print(f"[ONNX_MODEL] input={self.input_name}") + print(f"[ONNX_MODEL] outputs={self.output_names}") + + def _create_session( + self, + onnx_path: Path, + provider: str, + trt_home: Optional[str], + trt_fp16: bool, + ): + try: + import onnxruntime as ort + except ImportError: + raise ImportError( + "onnxruntime não está instalado. Use:\n" + " pip install onnxruntime-gpu" + ) + + provider = provider.lower() + + 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 que vocês estão usando + dll_dirs.append(r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v12.3\bin") + 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 == "cpu": + providers = ["CPUExecutionProvider"] + + elif provider == "cuda": + providers = ["CUDAExecutionProvider", "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), + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str(cache_dir), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(cache_dir), + "trt_max_workspace_size": 4 * 1024 * 1024 * 1024, + } + + providers = [ + ("TensorrtExecutionProvider", trt_options), + "CUDAExecutionProvider", + "CPUExecutionProvider", + ] + + else: + raise RuntimeError(f"Provider ONNX inválido: {provider}") + + 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 ONNX disponível. Pedido={providers}, disponíveis={available}") + + session = ort.InferenceSession( + str(onnx_path), + sess_options=sess_options, + providers=providers_ok, + ) + + active = session.get_providers() + print(f"[ONNX] usando providers: {active}") + + if provider == "tensorrt" and "TensorrtExecutionProvider" not in active: + raise RuntimeError(f"TensorRT solicitado, mas não ficou ativo. Providers ativos: {active}") + + if provider == "cuda" and "CUDAExecutionProvider" not in active: + raise RuntimeError(f"CUDA solicitado, mas não ficou ativo. Providers ativos: {active}") + + return session + + def _normalize(self, x: np.ndarray) -> np.ndarray: + if self.mean is not None and self.std is not None: + return ((x - self.mean[0]) / np.clip(self.std[0], 1e-6, None)).astype(np.float32) + return x.astype(np.float32) + + def _selected_heads(self) -> Optional[List[str]]: + if self.runtime_mode in ("target_direct", "target_head"): + return ["target"] + if self.runtime_mode in ("operational", "target_op"): + return ["vegetation", "cana"] + return None + + @staticmethod + 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) + + @staticmethod + def _resize_logits_nchw(logits: np.ndarray, target_hw: Tuple[int, int]) -> np.ndarray: + 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 _map_outputs(self, outputs: List[np.ndarray]) -> Dict[str, np.ndarray]: + raw = { + name: arr.astype(np.float32) + for name, arr in zip(self.output_names, outputs) + } + + mapped = {} + for head in self.heads_config.keys(): + candidates = [ + head, + f"{head}_logits", + f"output_{head}", + ] + + found = None + for c in candidates: + if c in raw: + found = c + break + + if found is not None: + mapped[head] = raw[found] + + return mapped + + def infer(self, chw_01: np.ndarray) -> Tuple[Dict[str, np.ndarray], Dict[str, np.ndarray], float]: + h, w = int(chw_01.shape[1]), int(chw_01.shape[2]) + + x = self._normalize(chw_01) + x = np.expand_dims(x, axis=0).astype(np.float32) + + t0 = time.perf_counter() + outputs = self.session.run(None, {self.input_name: x}) + t_ms = (time.perf_counter() - t0) * 1000.0 + + logits_by_head = self._map_outputs(outputs) + + selected = self._selected_heads() + if selected is not None: + logits_by_head = { + hname: logits + for hname, logits in logits_by_head.items() + if hname in selected + } + + preds = {} + probs = {} + + for head_name, logits in logits_by_head.items(): + logits = self._resize_logits_nchw(logits, (h, w)) + prob = self._softmax_np(logits, axis=1)[0] + pred = np.argmax(prob, axis=0).astype(np.uint8) + + preds[head_name] = pred + probs[head_name] = prob.astype(np.float32) + + return preds, probs, t_ms + + # ============================================================ # Checkpoints / paths # ============================================================ @@ -1084,6 +1319,43 @@ def find_checkpoint(save_dir: Path, preferred: Optional[str] = None) -> Path: raise FileNotFoundError("Nenhum checkpoint encontrado. Procurei:\n" + "\n".join(str(c) for c in candidates)) +def find_onnx_model(save_dir: Path, ckpt_path: Path, preferred: Optional[str] = None) -> Path: + """ + Resolve o .onnx. + + Se preferred for informado, usa ele. + Caso contrário, usa o mesmo stem do checkpoint: + best_score.pt -> best_score.onnx + """ + if preferred: + p = Path(preferred) + if not p.is_absolute(): + p_cwd = (Path.cwd() / p).resolve() + p_save = (save_dir / p).resolve() + p = p_cwd if p_cwd.is_file() else p_save + + if not p.is_file(): + raise FileNotFoundError(f"ONNX não encontrado: {p}") + + return p.resolve() + + p = ckpt_path.with_suffix(".onnx") + + if not p.is_file(): + alt = save_dir / f"{ckpt_path.stem}.onnx" + p = alt + + if not p.is_file(): + raise FileNotFoundError( + f"ONNX não encontrado para checkpoint {ckpt_path.name}. Procurei:\n" + f" {ckpt_path.with_suffix('.onnx')}\n" + f" {save_dir / (ckpt_path.stem + '.onnx')}\n" + f"Informe manualmente com --onnx." + ) + + return p.resolve() + + # ============================================================ # Visualização # ============================================================ @@ -1179,6 +1451,37 @@ def class_percent(mask: np.ndarray, class_id: int, ignore_id: int = 255) -> floa return float(((mask == class_id) & valid).sum() * 100.0 / den) +def compare_pred_equal_percent(a: Optional[np.ndarray], b: Optional[np.ndarray]) -> Optional[float]: + if a is None or b is None: + return None + + if a.shape != b.shape: + b = cv2.resize( + b.astype(np.uint8), + (a.shape[1], a.shape[0]), + interpolation=cv2.INTER_NEAREST, + ) + + return float(np.mean(a == b) * 100.0) + + +def diff_mask_rgb(a: Optional[np.ndarray], b: Optional[np.ndarray]) -> Optional[np.ndarray]: + if a is None or b is None: + return None + + if a.shape != b.shape: + b = cv2.resize( + b.astype(np.uint8), + (a.shape[1], a.shape[0]), + interpolation=cv2.INTER_NEAREST, + ) + + diff = (a != b).astype(np.uint8) * 255 + rgb = np.zeros((diff.shape[0], diff.shape[1], 3), dtype=np.uint8) + rgb[:, :, 0] = diff # vermelho em RGB + return rgb + + # ============================================================ # Métricas numpy # ============================================================ @@ -1240,6 +1543,10 @@ def main(): parser.add_argument("--start_idx", type=int, default=0) parser.add_argument("--max_width", type=int, default=1800) parser.add_argument("--runtime_mode", default="all", choices=["all", "target_direct", "operational"]) + parser.add_argument("--onnx", default="") + parser.add_argument("--onnx_provider", default="", choices=["", "cpu", "cuda", "tensorrt"]) + parser.add_argument("--trt_home", default=None) + parser.add_argument("--trt_no_fp16", action="store_true") args = parser.parse_args() config_path = resolve_path(args.config, Path.cwd()) @@ -1336,6 +1643,32 @@ def main(): use_amp=not args.no_amp, ) + onnx_tester = None + onnx_path = None + + if args.onnx_provider: + onnx_path = find_onnx_model( + save_dir=save_dir, + ckpt_path=ckpt_path, + preferred=args.onnx, + ) + + onnx_tester = OnnxMultiHeadTester( + config=config, + onnx_path=onnx_path, + provider=args.onnx_provider, + channels=channels, + heads_config=heads_config, + mean=mean, + std=std, + trt_home=args.trt_home, + trt_fp16=not args.trt_no_fp16, + ) + + if onnx_tester is not None: + print(f"ONNX : {onnx_path}") + print(f"ONNX provider: {args.onnx_provider}") + out_dir = Path(args.out_dir) ensure_dir(out_dir) @@ -1351,6 +1684,10 @@ def main(): win_name = "OAK-FCC-3 MultiHead Test | D/A navega | S salva | SPACE detalhado | Q sai" cv2.namedWindow(win_name, cv2.WINDOW_NORMAL) + win_name_onnx = None + if onnx_tester is not None: + win_name_onnx = "ONNX/TensorRT MultiHead Test | comparação visual" + cv2.namedWindow(win_name_onnx, cv2.WINDOW_NORMAL) while True: sample = samples[idx] @@ -1368,6 +1705,13 @@ def main(): preds, probs, t_inf = tester.infer(chw) preview_rgb = tensor_to_preview_rgb(chw) + onnx_preds = None + onnx_probs = None + t_onnx = None + + if onnx_tester is not None: + onnx_preds, onnx_probs, t_onnx = onnx_tester.infer(chw) + first_pred = next(iter(preds.values())) pred_h, pred_w = first_pred.shape[:2] @@ -1508,7 +1852,143 @@ def main(): cv2.putText(header, h2, (12, 58), cv2.FONT_HERSHEY_SIMPLEX, 0.52, (0, 255, 120), 1, cv2.LINE_AA) canvas = np.vstack([header, canvas]) + onnx_canvas = None + if onnx_tester is not None and onnx_preds is not None and onnx_probs is not None: + onnx_pred_sem = onnx_preds.get("semantic") + onnx_pred_veg = onnx_preds.get("vegetation") + onnx_pred_cana = onnx_preds.get("cana") + + onnx_pred_target_op = None + if onnx_pred_veg is not None and onnx_pred_cana is not None: + onnx_pred_target_op = operational_target_mask( + onnx_pred_veg, + onnx_pred_cana, + ignore_id=ignore_id, + ) + + onnx_pred_target_head = onnx_preds.get("target") + onnx_pred_target = onnx_pred_target_head if onnx_pred_target_head is not None else onnx_pred_target_op + + onnx_prob_veg = onnx_probs["vegetation"][1] if "vegetation" in onnx_probs and onnx_probs["vegetation"].shape[0] > 1 else None + onnx_prob_cana = onnx_probs["cana"][1] if "cana" in onnx_probs and onnx_probs["cana"].shape[0] > 1 else None + + onnx_prob_target_op = None + if onnx_prob_veg is not None and onnx_prob_cana is not None: + onnx_prob_target_op = np.clip(onnx_prob_veg * (1.0 - onnx_prob_cana), 0.0, 1.0) + + onnx_prob_target_head = None + if "target" in onnx_probs: + onnx_prob_target_head = onnx_probs["target"][1] if onnx_probs["target"].shape[0] > 1 else onnx_probs["target"][0] + + onnx_prob_target = onnx_prob_target_head if onnx_prob_target_head is not None else onnx_prob_target_op + + eq_sem = compare_pred_equal_percent(pred_sem, onnx_pred_sem) + eq_veg = compare_pred_equal_percent(pred_veg, onnx_pred_veg) + eq_cana = compare_pred_equal_percent(pred_cana, onnx_pred_cana) + eq_target = compare_pred_equal_percent(pred_target, onnx_pred_target) + + onnx_panels: List[Tuple[str, np.ndarray, str]] = [] + + if onnx_pred_sem is not None: + onnx_sem_rgb = ids_to_rgb(onnx_pred_sem, semantic_cmap, ignore_id) + onnx_panels.append(( + "ONNX semantic", + onnx_sem_rgb, + "" if eq_sem is None else f"igual PT={eq_sem:.3f}%" + )) + onnx_panels.append(( + "ONNX overlay semantic", + overlay_rgb(preview_rgb, onnx_sem_rgb, args.alpha), + "" + )) + + d = diff_mask_rgb(pred_sem, onnx_pred_sem) + if d is not None: + onnx_panels.append(("Diff semantic", d, "vermelho=diferente")) + + if onnx_pred_veg is not None: + onnx_veg_rgb = ids_to_rgb(onnx_pred_veg, BINARY_COLORS_RGB, ignore_id) + onnx_panels.append(( + "ONNX vegetation", + onnx_veg_rgb, + "" if eq_veg is None else f"igual PT={eq_veg:.3f}%" + )) + + if onnx_prob_veg is not None: + onnx_panels.append(( + "ONNX P vegetation", + prob_to_heat_rgb(onnx_prob_veg), + f"mean={float(onnx_prob_veg.mean()):.3f}" + )) + + if onnx_pred_cana is not None: + onnx_cana_rgb = ids_to_rgb(onnx_pred_cana, CANA_COLORS_RGB, ignore_id) + onnx_panels.append(( + "ONNX cana", + onnx_cana_rgb, + "" if eq_cana is None else f"igual PT={eq_cana:.3f}%" + )) + + if onnx_prob_cana is not None: + onnx_panels.append(( + "ONNX P cana", + prob_to_heat_rgb(onnx_prob_cana), + f"mean={float(onnx_prob_cana.mean()):.3f}" + )) + + if onnx_pred_target is not None: + onnx_target_rgb = ids_to_rgb(onnx_pred_target, TARGET_COLORS_RGB, ignore_id) + title = "ONNX target HEAD" if onnx_pred_target_head is not None else "ONNX target OP" + + onnx_panels.append(( + title, + onnx_target_rgb, + "" if eq_target is None else f"igual PT={eq_target:.3f}%" + )) + onnx_panels.append(( + "ONNX overlay target", + overlay_rgb(preview_rgb, onnx_target_rgb, args.alpha), + f"inf={t_onnx:.1f}ms" + )) + + if onnx_prob_target is not None: + onnx_panels.append(( + "ONNX P target", + prob_to_heat_rgb(onnx_prob_target), + f"mean={float(onnx_prob_target.mean()):.3f}" + )) + + d = diff_mask_rgb(pred_target, onnx_pred_target) + if d is not None: + onnx_panels.append(("Diff target", d, "vermelho=diferente")) + + onnx_canvas = compose_grid(onnx_panels, cols=3, max_width=args.max_width) + + header_h_onnx = 78 + header_onnx = np.zeros((header_h_onnx, onnx_canvas.shape[1], 3), dtype=np.uint8) + header_onnx[:] = (18, 18, 35) + + h1_onnx = f"ONNX {args.onnx_provider} | {source_name} | inf={t_onnx:.1f}ms" + h2_parts = [] + if eq_sem is not None: + h2_parts.append(f"sem={eq_sem:.3f}%") + if eq_veg is not None: + h2_parts.append(f"veg={eq_veg:.3f}%") + if eq_cana is not None: + h2_parts.append(f"cana={eq_cana:.3f}%") + if eq_target is not None: + h2_parts.append(f"target={eq_target:.3f}%") + h2_onnx = "igual PyTorch: " + " | ".join(h2_parts) if h2_parts else "comparação indisponível" + + cv2.putText(header_onnx, h1_onnx, (12, 28), cv2.FONT_HERSHEY_SIMPLEX, 0.65, (235, 235, 235), 1, cv2.LINE_AA) + cv2.putText(header_onnx, h2_onnx, (12, 58), cv2.FONT_HERSHEY_SIMPLEX, 0.52, (0, 255, 120), 1, cv2.LINE_AA) + + onnx_canvas = np.vstack([header_onnx, onnx_canvas]) + cv2.imshow(win_name, cv2.cvtColor(canvas, cv2.COLOR_RGB2BGR)) + if onnx_canvas is not None and win_name_onnx is not None: + cv2.imshow(win_name_onnx, cv2.cvtColor(onnx_canvas, cv2.COLOR_RGB2BGR)) + k = cv2.waitKey(0) & 0xFF if k in (ord("q"), ord("Q"), 27): @@ -1520,10 +2000,15 @@ def main(): elif k == ord(" "): detailed = not detailed elif k in (ord("s"), ord("S")): - out_path = out_dir / f"multihead_{idx:05d}_{sample.base}.png" + out_path = out_dir / f"multihead_pytorch_{idx:05d}_{sample.base}.png" cv2.imwrite(str(out_path), cv2.cvtColor(canvas, cv2.COLOR_RGB2BGR)) print(f"[SAVE] {out_path}") + if onnx_canvas is not None: + out_path_onnx = out_dir / f"multihead_onnx_{args.onnx_provider}_{idx:05d}_{sample.base}.png" + cv2.imwrite(str(out_path_onnx), cv2.cvtColor(onnx_canvas, cv2.COLOR_RGB2BGR)) + print(f"[SAVE] {out_path_onnx}") + cv2.destroyAllWindows() if visited: diff --git a/Python/OAK/datasets/oak-fcc-3/config.json b/Python/OAK/datasets/oak-fcc-3/config.json index 4e8a302cc..8c2bd90ad 100644 --- a/Python/OAK/datasets/oak-fcc-3/config.json +++ b/Python/OAK/datasets/oak-fcc-3/config.json @@ -1,7 +1,7 @@ { "camera": "oak-fcc-3", "modelo": "segformer_b1", - "model_name": "target_teached", + "model_name": "target_aug", "main_class_name": "cana", "es_classes": "", "model_to_use": "geral", @@ -15,8 +15,9 @@ "use_ndvi": false, "backbone": "nvidia/mit-b1", "fusion_mode": "stacked", - "stats_source_tag": "stacked_raw4", + "stats_source_tag": "stacked_raw5", "module_params_json": "calibration/module_params.json", + "ckpt_test": "best_target", "multi_head": true, "heads": { "semantic": {