#!/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()