From 760d9960605e8e7ca4f7b17f9f7bff97d5583e03 Mon Sep 17 00:00:00 2001 From: Diego Freitas Date: Thu, 17 Sep 2026 13:58:42 -0300 Subject: [PATCH] arquivos restantes do visual worker --- .../visual_worker/inferencia/__init__.py | 0 .../inferencia/corridor_onnx_contract.py | 316 +++++++++++ .../inferencia/corridor_onnx_runner.py | 512 ++++++++++++++++++ 3 files changed, 828 insertions(+) create mode 100644 AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/__init__.py create mode 100644 AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_contract.py create mode 100644 AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_runner.py diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/__init__.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_contract.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_contract.py new file mode 100644 index 000000000..23b5ae058 --- /dev/null +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_contract.py @@ -0,0 +1,316 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Dict, Mapping, Sequence, Tuple + + +EXPECTED_SCHEMA = "agrobot.visual.corridor.onnx.v2" +EXPECTED_CONTRACT_VERSION = 2 + + +def _int_key_dict(value: Any) -> Dict[int, str]: + if not isinstance(value, Mapping): + return {} + out: Dict[int, str] = {} + for key, name in value.items(): + out[int(key)] = str(name) + return out + + +def _fixed_shape(value: Sequence[Any], name: str) -> Tuple[int, ...]: + if not isinstance(value, (list, tuple)): + raise RuntimeError(f"{name} precisa ser lista/tupla, veio {value!r}") + + out = [] + for item in value: + if isinstance(item, bool) or not isinstance(item, int): + raise RuntimeError( + f"{name} precisa ter shape fixo inteiro; veio {list(value)!r}" + ) + out.append(int(item)) + return tuple(out) + + +def _ort_type_matches(actual: str, expected: str) -> bool: + actual = str(actual).strip().lower() + expected = str(expected).strip().lower() + aliases = { + "uint8": {"tensor(uint8)", "uint8"}, + "int32": {"tensor(int32)", "int32"}, + "float32": {"tensor(float)", "tensor(float32)", "float", "float32"}, + } + return actual in aliases.get(expected, {expected}) + + +@dataclass(frozen=True) +class CorridorOnnxContract: + schema: str + contract_version: int + + width: int + height: int + roi_inicio: float + roi_tamanho: float + + input_name: str + input_dtype: str + input_layout: str + input_color_order: str + input_shape: Tuple[int, ...] + + seg_output_name: str + seg_output_dtype: str + seg_output_shape: Tuple[int, ...] + + status_output_name: str + status_output_dtype: str + status_output_shape: Tuple[int, ...] + + semantic_id2label: Dict[int, str] + status_id2label: Dict[int, str] + nav_class_id: int + non_nav_class_id: int + + exporter_version: str + trainer_version: str + checkpoint_sha256: str + raw_contract: Dict[str, Any] + + @property + def resolution_wh(self) -> Tuple[int, int]: + return self.width, self.height + + @property + def num_status_classes(self) -> int: + return len(self.status_id2label) + + @property + def semantic_classes(self): + return [self.semantic_id2label[i] for i in sorted(self.semantic_id2label)] + + @property + def status_classes(self): + return [self.status_id2label[i] for i in sorted(self.status_id2label)] + + @classmethod + def from_session(cls, session) -> "CorridorOnnxContract": + meta = session.get_modelmeta() + custom = dict(getattr(meta, "custom_metadata_map", {}) or {}) + + contract_text = custom.get("agrobot.contract_json") + if not contract_text: + raise RuntimeError( + "ONNX recusado: metadata 'agrobot.contract_json' ausente. " + "O Visual Worker v2 aceita somente o contrato oficial de campo." + ) + + try: + raw = json.loads(contract_text) + except Exception as exc: + raise RuntimeError( + f"ONNX recusado: agrobot.contract_json inválido: {exc}" + ) from exc + + schema = str(raw.get("schema", "")) + version = int(raw.get("contract_version", -1)) + + if schema != EXPECTED_SCHEMA: + raise RuntimeError( + f"ONNX recusado: schema={schema!r}; esperado={EXPECTED_SCHEMA!r}." + ) + if version != EXPECTED_CONTRACT_VERSION: + raise RuntimeError( + f"ONNX recusado: contract_version={version}; " + f"esperado={EXPECTED_CONTRACT_VERSION}." + ) + + inp = raw.get("input") or {} + outputs = raw.get("outputs") or {} + seg = outputs.get("seg_ids") or {} + status = outputs.get("label_probs") or {} + geometry = inp.get("geometry") or {} + model = raw.get("model") or {} + + input_shape = _fixed_shape(inp.get("shape", []), "input.shape") + seg_shape = _fixed_shape(seg.get("shape", []), "seg_ids.shape") + status_shape = _fixed_shape(status.get("shape", []), "label_probs.shape") + + if len(input_shape) != 4 or input_shape[0] != 1 or input_shape[-1] != 3: + raise RuntimeError(f"Contrato input shape inválido: {input_shape}") + + height = int(input_shape[1]) + width = int(input_shape[2]) + + expected_seg = (1, height, width) + if seg_shape != expected_seg: + raise RuntimeError( + f"Contrato seg_ids shape inválido: {seg_shape}; esperado={expected_seg}" + ) + + semantic_id2label = _int_key_dict(seg.get("id2label")) + status_id2label = _int_key_dict(status.get("id2label")) + if not semantic_id2label: + raise RuntimeError("Contrato sem semantic id2label.") + if not status_id2label: + raise RuntimeError("Contrato sem status id2label.") + + expected_status = (1, len(status_id2label)) + if status_shape != expected_status: + raise RuntimeError( + f"Contrato label_probs shape inválido: {status_shape}; " + f"esperado={expected_status}" + ) + + nav_id = int(seg.get("nav_class_id")) + non_nav_id = int(seg.get("non_nav_class_id")) + if nav_id == non_nav_id: + raise RuntimeError("nav_class_id e non_nav_class_id são iguais.") + if nav_id not in semantic_id2label or non_nav_id not in semantic_id2label: + raise RuntimeError( + f"IDs semânticos inválidos: nav={nav_id} non_nav={non_nav_id} " + f"classes={semantic_id2label}" + ) + + contract = cls( + schema=schema, + contract_version=version, + width=width, + height=height, + roi_inicio=float(geometry.get("source_roi_inicio", 0.0)), + roi_tamanho=float(geometry.get("source_roi_tamanho", 1.0)), + input_name=str(inp.get("name", "")), + input_dtype=str(inp.get("dtype", "")), + input_layout=str(inp.get("layout", "")).upper(), + input_color_order=str(inp.get("color_order", "")).upper(), + input_shape=input_shape, + seg_output_name=str(seg.get("name", "")), + seg_output_dtype=str(seg.get("dtype", "")), + seg_output_shape=seg_shape, + status_output_name=str(status.get("name", "")), + status_output_dtype=str(status.get("dtype", "")), + status_output_shape=status_shape, + semantic_id2label=semantic_id2label, + status_id2label=status_id2label, + nav_class_id=nav_id, + non_nav_class_id=non_nav_id, + exporter_version=str(raw.get("exporter_version", "unknown")), + trainer_version=str(model.get("trainer_version", "unknown")), + checkpoint_sha256=str(model.get("checkpoint_sha256", "")), + raw_contract=raw, + ) + + contract.validate_session(session) + return contract + + def validate_session(self, session) -> None: + if self.input_name != "rgb_u8_nhwc": + raise RuntimeError( + f"Input name inválido: {self.input_name!r}; esperado 'rgb_u8_nhwc'." + ) + if self.input_dtype.lower() != "uint8": + raise RuntimeError(f"Input dtype inválido no contrato: {self.input_dtype}") + if self.input_layout != "NHWC": + raise RuntimeError(f"Input layout inválido: {self.input_layout}") + if self.input_color_order != "RGB": + raise RuntimeError(f"Input color_order inválido: {self.input_color_order}") + if self.seg_output_name != "seg_ids": + raise RuntimeError(f"Seg output inválido: {self.seg_output_name}") + if self.status_output_name != "label_probs": + raise RuntimeError(f"Status output inválido: {self.status_output_name}") + if self.seg_output_dtype.lower() != "int32": + raise RuntimeError(f"seg_ids dtype inválido: {self.seg_output_dtype}") + if self.status_output_dtype.lower() != "float32": + raise RuntimeError(f"label_probs dtype inválido: {self.status_output_dtype}") + + inputs = {x.name: x for x in session.get_inputs()} + outputs = {x.name: x for x in session.get_outputs()} + + if set(inputs) != {self.input_name}: + raise RuntimeError( + f"ONNX precisa expor somente input {self.input_name!r}; recebeu={list(inputs)}" + ) + if set(outputs) != {self.seg_output_name, self.status_output_name}: + raise RuntimeError( + "ONNX outputs divergentes do contrato: " + f"recebeu={list(outputs)}" + ) + + inp = inputs[self.input_name] + seg = outputs[self.seg_output_name] + status = outputs[self.status_output_name] + + if not _ort_type_matches(inp.type, self.input_dtype): + raise RuntimeError(f"ORT input dtype={inp.type}; contrato={self.input_dtype}") + if not _ort_type_matches(seg.type, self.seg_output_dtype): + raise RuntimeError(f"ORT seg dtype={seg.type}; contrato={self.seg_output_dtype}") + if not _ort_type_matches(status.type, self.status_output_dtype): + raise RuntimeError( + f"ORT status dtype={status.type}; contrato={self.status_output_dtype}" + ) + + if _fixed_shape(inp.shape, "ORT input shape") != self.input_shape: + raise RuntimeError( + f"ORT input shape={inp.shape}; contrato={self.input_shape}" + ) + if _fixed_shape(seg.shape, "ORT seg shape") != self.seg_output_shape: + raise RuntimeError( + f"ORT seg shape={seg.shape}; contrato={self.seg_output_shape}" + ) + if _fixed_shape(status.shape, "ORT status shape") != self.status_output_shape: + raise RuntimeError( + f"ORT status shape={status.shape}; contrato={self.status_output_shape}" + ) + + def validate_status_enum(self, enum_type) -> None: + missing = [] + for _idx, name in sorted(self.status_id2label.items()): + try: + enum_type[str(name)] + except Exception: + missing.append(str(name)) + + if missing: + raise RuntimeError( + "Status do ONNX não existem em StatusCarroMapa: " + f"{missing}. O runtime mapeia por NOME, nunca pela posição do ID." + ) + + def status_name(self, label_id: int) -> str: + return self.status_id2label.get(int(label_id), f"status_{int(label_id)}") + + def semantic_name(self, class_id: int) -> str: + return self.semantic_id2label.get(int(class_id), f"class_{int(class_id)}") + + def runtime_summary(self) -> Dict[str, Any]: + return { + "schema": self.schema, + "contract_version": self.contract_version, + "resolution_wh": [self.width, self.height], + "roi_inicio": self.roi_inicio, + "roi_tamanho": self.roi_tamanho, + "input": { + "name": self.input_name, + "dtype": self.input_dtype, + "layout": self.input_layout, + "color_order": self.input_color_order, + "shape": list(self.input_shape), + }, + "outputs": { + self.seg_output_name: { + "dtype": self.seg_output_dtype, + "shape": list(self.seg_output_shape), + }, + self.status_output_name: { + "dtype": self.status_output_dtype, + "shape": list(self.status_output_shape), + }, + }, + "semantic_id2label": dict(self.semantic_id2label), + "status_id2label": dict(self.status_id2label), + "nav_class_id": self.nav_class_id, + "non_nav_class_id": self.non_nav_class_id, + "exporter_version": self.exporter_version, + "trainer_version": self.trainer_version, + "checkpoint_sha256": self.checkpoint_sha256, + } diff --git a/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_runner.py b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_runner.py new file mode 100644 index 000000000..ee887233f --- /dev/null +++ b/AgroBase/AgroBase/bin/x64/Debug/Python/Scripts/workers/visual_worker/inferencia/corridor_onnx_runner.py @@ -0,0 +1,512 @@ +from __future__ import annotations + +import hashlib +import json +import math +import os +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import cv2 +import numpy as np +import onnxruntime as ort + +from visual_worker.inferencia.corridor_onnx_contract import CorridorOnnxContract + + +@dataclass +class CorridorInferenceResult: + seg_ids: np.ndarray + status: Dict[str, Any] + timing_ms: Dict[str, float] + result_ts: float + + +class CorridorOnnxRunner: + """ + Runtime oficial de campo do corredor. + + ÚNICO contrato aceito: + input : rgb_u8_nhwc uint8 [1,H,W,3] RGB + output: seg_ids int32 [1,H,W] + output: label_probs float32 [1,K] + + O grafo já contém: + cast -> /255 -> NHWC/NCHW -> mean/std -> SegFormer -> + resize/argmax -> StatusHead/softmax. + + Fora do ONNX ficam somente geometria da câmera e conversão de cor da fonte. + """ + + def __init__( + self, + onnx_path: str | os.PathLike, + runtime_config: Optional[Dict[str, Any]] = None, + telemetry_config: Optional[Dict[str, Any]] = None, + mostrar_log=None, + ): + self.onnx_path = Path(onnx_path).resolve() + self.runtime_config = dict(runtime_config or {}) + self.telemetry_config = dict(telemetry_config or {}) + self.mostrar_log = mostrar_log or print + + if not self.onnx_path.is_file(): + raise FileNotFoundError(f"ONNX de corredor não encontrado: {self.onnx_path}") + + self.provider_requested = str( + self.runtime_config.get("provider", "tensorrt") + ).strip().lower() + if self.provider_requested == "trt": + self.provider_requested = "tensorrt" + + if self.provider_requested not in {"tensorrt", "cuda", "cpu"}: + raise ValueError( + f"model_runtime.provider inválido: {self.provider_requested}" + ) + + self.allow_fallback = bool( + self.runtime_config.get("allow_provider_fallback", True) + ) + self.camera_frame_color = str( + self.runtime_config.get("camera_frame_color", "BGR") + ).strip().upper() + if self.camera_frame_color not in {"BGR", "RGB"}: + raise ValueError( + f"camera_frame_color deve ser BGR ou RGB; veio {self.camera_frame_color}" + ) + + self.validate_outputs_each_inference = bool( + self.runtime_config.get("validate_outputs_each_inference", False) + ) + self.include_status_probs = bool( + self.telemetry_config.get("include_status_probs", False) + ) + self.detailed_timing = bool( + self.telemetry_config.get("runner_detailed_timing", True) + ) + self.log_startup = bool( + self.telemetry_config.get("log_runtime_startup", True) + ) + + self.model_sha256 = self._sha256_file(self.onnx_path) + self.model_sha_short = self.model_sha256[:16] + self.trt_cache_dir = self._resolve_trt_cache_dir() + + self._verify_sidecar_if_requested() + + t0 = time.perf_counter() + self.session = self._create_session() + self.session_create_ms = (time.perf_counter() - t0) * 1000.0 + + self.contract = CorridorOnnxContract.from_session(self.session) + self.input_name = self.contract.input_name + self.output_names = [ + self.contract.seg_output_name, + self.contract.status_output_name, + ] + + self.provider_chain = list(self.session.get_providers()) + self.provider_primary = self.provider_chain[0] if self.provider_chain else "unknown" + self.provider_expected_ort = self._provider_mode_to_ort(self.provider_requested) + self.runtime_degraded = self.provider_primary != self.provider_expected_ort + + if self.runtime_degraded and not self.allow_fallback: + raise RuntimeError( + "Provider solicitado não ficou primário e fallback está proibido: " + f"requested={self.provider_expected_ort} chain={self.provider_chain}" + ) + + # Preview/debug apenas. O modelo não depende de labelmap externo. + self.classes = self.contract.semantic_classes + self.colormap_rgb = self._default_semantic_colormap_rgb() + + self.inference_count = 0 + self.last_infer_ts = 0.0 + self.last_timing_ms: Dict[str, float] = {} + self.first_inference_ms: Optional[float] = None + + if self.log_startup: + self._log_runtime_summary() + + # ------------------------------------------------------------------ + # Session / provider + # ------------------------------------------------------------------ + + def _create_session(self) -> ort.InferenceSession: + options = ort.SessionOptions() + options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL + + intra = int(self.runtime_config.get("ort_intra_op_num_threads", 1)) + inter = int(self.runtime_config.get("ort_inter_op_num_threads", 1)) + if intra > 0: + options.intra_op_num_threads = intra + if inter > 0: + options.inter_op_num_threads = inter + + providers = self._build_providers() + + try: + return ort.InferenceSession( + str(self.onnx_path), + sess_options=options, + providers=providers, + ) + except Exception as exc: + raise RuntimeError( + f"Falha ao abrir ONNX v2 {self.onnx_path} | " + f"provider={self.provider_requested}: {exc}" + ) from exc + + def _build_providers(self) -> List[Any]: + available = set(ort.get_available_providers()) + requested_ort = self._provider_mode_to_ort(self.provider_requested) + + if requested_ort not in available and not self.allow_fallback: + raise RuntimeError( + f"Provider obrigatório {requested_ort} indisponível. " + f"Disponíveis={sorted(available)}" + ) + + providers: List[Any] = [] + + if self.provider_requested == "tensorrt": + if "TensorrtExecutionProvider" in available: + self.trt_cache_dir.mkdir(parents=True, exist_ok=True) + trt_options: Dict[str, Any] = { + "trt_fp16_enable": bool( + self.runtime_config.get("trt_fp16_enable", True) + ), + "trt_engine_cache_enable": bool( + self.runtime_config.get("trt_engine_cache_enable", True) + ), + "trt_engine_cache_path": str(self.trt_cache_dir), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(self.trt_cache_dir), + } + + workspace = self.runtime_config.get("trt_max_workspace_size") + if workspace is not None: + trt_options["trt_max_workspace_size"] = int(workspace) + + providers.append(("TensorrtExecutionProvider", trt_options)) + + if self.allow_fallback and "CUDAExecutionProvider" in available: + providers.append("CUDAExecutionProvider") + if self.allow_fallback and "CPUExecutionProvider" in available: + providers.append("CPUExecutionProvider") + + elif self.provider_requested == "cuda": + if "CUDAExecutionProvider" in available: + providers.append("CUDAExecutionProvider") + if self.allow_fallback and "CPUExecutionProvider" in available: + providers.append("CPUExecutionProvider") + + else: + if "CPUExecutionProvider" in available: + providers.append("CPUExecutionProvider") + + if not providers: + raise RuntimeError( + f"Nenhum provider utilizável. requested={self.provider_requested} " + f"available={sorted(available)}" + ) + + return providers + + @staticmethod + def _provider_mode_to_ort(mode: str) -> str: + return { + "tensorrt": "TensorrtExecutionProvider", + "cuda": "CUDAExecutionProvider", + "cpu": "CPUExecutionProvider", + }[mode] + + def _resolve_trt_cache_dir(self) -> Path: + root = Path( + self.runtime_config.get("trt_engine_cache_root", "./trt_cache_visual_worker") + ) + if not root.is_absolute(): + root = Path.cwd() / root + # Cada modelo ganha sua própria gaveta. Engine antiga nunca cruza modelo novo. + return (root / self.model_sha_short).resolve() + + # ------------------------------------------------------------------ + # Contract / sidecar + # ------------------------------------------------------------------ + + def _verify_sidecar_if_requested(self) -> None: + if not bool(self.runtime_config.get("verify_sidecar_sha256", True)): + return + + sidecar = self.onnx_path.with_suffix(".contract.json") + if not sidecar.is_file(): + return + + try: + data = json.loads(sidecar.read_text(encoding="utf-8")) + expected = str((data.get("artifact") or {}).get("onnx_sha256") or "") + except Exception as exc: + raise RuntimeError(f"Sidecar inválido {sidecar}: {exc}") from exc + + if expected and expected.lower() != self.model_sha256.lower(): + raise RuntimeError( + "SHA256 do ONNX diverge do sidecar: " + f"arquivo={self.model_sha256} sidecar={expected}" + ) + + # ------------------------------------------------------------------ + # Inference + # ------------------------------------------------------------------ + + def infer(self, frame_camera: np.ndarray) -> CorridorInferenceResult: + t0 = time.perf_counter() + + frame = self._validate_source_frame(frame_camera) + t_valid = time.perf_counter() + + input_np, prep_timing = self._prepare_contract_input(frame) + t_prepared = time.perf_counter() + + outputs = self.session.run( + self.output_names, + {self.input_name: input_np}, + ) + t_session = time.perf_counter() + + seg_ids = self._decode_seg_ids(outputs[0]) + status = self._decode_status_probs(outputs[1]) + t_decode = time.perf_counter() + + timing = { + "validate_ms": (t_valid - t0) * 1000.0, + **prep_timing, + "prepare_total_ms": (t_prepared - t_valid) * 1000.0, + "session_ms": (t_session - t_prepared) * 1000.0, + "decode_ms": (t_decode - t_session) * 1000.0, + "total_ms": (t_decode - t0) * 1000.0, + } + + now = time.time() + self.inference_count += 1 + self.last_infer_ts = now + self.last_timing_ms = timing + if self.first_inference_ms is None: + self.first_inference_ms = float(timing["total_ms"]) + + return CorridorInferenceResult( + seg_ids=seg_ids, + status=status, + timing_ms=timing, + result_ts=now, + ) + + def _prepare_contract_input(self, frame: np.ndarray) -> Tuple[np.ndarray, Dict[str, float]]: + t0 = time.perf_counter() + + h, w = frame.shape[:2] + y0, y1 = self._compute_roi_indices_top_origin( + h, + self.contract.roi_inicio, + self.contract.roi_tamanho, + ) + roi = frame[y0:y1, :, :3] + t_crop = time.perf_counter() + + target_w, target_h = self.contract.resolution_wh + if roi.shape[1] == target_w and roi.shape[0] == target_h: + resized = roi + else: + resized = cv2.resize( + roi, + (target_w, target_h), + interpolation=cv2.INTER_AREA, + ) + t_resize = time.perf_counter() + + if self.camera_frame_color == self.contract.input_color_order: + rgb = resized + elif self.camera_frame_color == "BGR" and self.contract.input_color_order == "RGB": + rgb = cv2.cvtColor(resized, cv2.COLOR_BGR2RGB) + elif self.camera_frame_color == "RGB" and self.contract.input_color_order == "BGR": + rgb = cv2.cvtColor(resized, cv2.COLOR_RGB2BGR) + else: + raise RuntimeError( + f"Conversão de cor não suportada: camera={self.camera_frame_color} " + f"model={self.contract.input_color_order}" + ) + t_color = time.perf_counter() + + inp = np.ascontiguousarray(rgb[None, ...], dtype=np.uint8) + t_contig = time.perf_counter() + + timing = { + "crop_ms": (t_crop - t0) * 1000.0, + "resize_ms": (t_resize - t_crop) * 1000.0, + "color_ms": (t_color - t_resize) * 1000.0, + "contiguous_ms": (t_contig - t_color) * 1000.0, + } + return inp, timing + + def _decode_seg_ids(self, output: np.ndarray) -> np.ndarray: + arr = np.asarray(output) + if arr.shape != self.contract.seg_output_shape: + raise RuntimeError( + f"seg_ids shape={arr.shape}; esperado={self.contract.seg_output_shape}" + ) + if arr.dtype != np.int32: + raise RuntimeError(f"seg_ids dtype={arr.dtype}; esperado=int32") + + seg = arr[0] + if self.validate_outputs_each_inference: + ids = np.unique(seg) + invalid = [ + int(x) for x in ids + if int(x) not in self.contract.semantic_id2label + ] + if invalid: + raise RuntimeError(f"seg_ids contém classes fora do contrato: {invalid}") + + return np.ascontiguousarray(seg) + + def _decode_status_probs(self, output: np.ndarray) -> Dict[str, Any]: + arr = np.asarray(output) + if arr.shape != self.contract.status_output_shape: + raise RuntimeError( + f"label_probs shape={arr.shape}; esperado={self.contract.status_output_shape}" + ) + if arr.dtype != np.float32: + arr = arr.astype(np.float32, copy=False) + + probs = np.ascontiguousarray(arr[0]) + if not np.all(np.isfinite(probs)): + raise RuntimeError("label_probs contém NaN/Inf.") + + total = float(probs.sum()) + if not (0.97 <= total <= 1.03): + raise RuntimeError( + f"label_probs não parecem softmax: soma={total:.6f}" + ) + + order = np.argsort(probs)[::-1] + top1 = int(order[0]) + top2 = int(order[1]) if len(order) > 1 else top1 + conf1 = float(probs[top1]) + conf2 = float(probs[top2]) if len(order) > 1 else 0.0 + margin = float(conf1 - conf2) + + eps = 1e-12 + entropy = float(-np.sum(probs * np.log(np.clip(probs, eps, 1.0)))) + max_entropy = math.log(max(len(probs), 2)) + entropy_norm = float(entropy / max_entropy) if max_entropy > 0 else 0.0 + + result: Dict[str, Any] = { + "label_id": top1, + "label_name": self.contract.status_name(top1), + "confidence": conf1, + "second_id": top2, + "second_name": self.contract.status_name(top2), + "second_confidence": conf2, + "margin": margin, + "entropy": entropy_norm, + } + + # Vetor completo é minúsculo, mas só entra no payload de debug se pedido. + result["probs_array"] = probs + if self.include_status_probs: + result["probs"] = probs.astype(float).tolist() + + return result + + @staticmethod + def _validate_source_frame(frame: np.ndarray) -> np.ndarray: + if frame is None or not hasattr(frame, "shape") or frame.size == 0: + raise ValueError("Frame da câmera vazio.") + if frame.ndim != 3 or frame.shape[2] != 3: + raise ValueError(f"Frame precisa ser HxWx3, veio {frame.shape}") + if frame.dtype != np.uint8: + frame = np.clip(frame, 0, 255).astype(np.uint8) + return frame + + @staticmethod + def _compute_roi_indices_top_origin( + height: int, + roi_inicio: float, + roi_tamanho: float, + ) -> Tuple[int, int]: + start = float(np.clip(roi_inicio, 0.0, 1.0)) + size = float(np.clip(roi_tamanho, 0.0, 1.0)) + end = min(1.0, start + size) + y0 = int(round(height * start)) + y1 = int(round(height * end)) + y0 = max(0, min(height - 1, y0)) + y1 = max(y0 + 1, min(height, y1)) + return y0, y1 + + # ------------------------------------------------------------------ + # Telemetry + # ------------------------------------------------------------------ + + def runtime_info(self) -> Dict[str, Any]: + return { + "model_path": str(self.onnx_path), + "model_sha256": self.model_sha256, + "model_sha_short": self.model_sha_short, + "provider_requested": self.provider_requested, + "provider_expected_ort": self.provider_expected_ort, + "provider_primary": self.provider_primary, + "provider_chain": list(self.provider_chain), + "runtime_degraded": bool(self.runtime_degraded), + "allow_provider_fallback": bool(self.allow_fallback), + "camera_frame_color": self.camera_frame_color, + "model_input_color": self.contract.input_color_order, + "session_create_ms": float(self.session_create_ms), + "first_inference_ms": self.first_inference_ms, + "inference_count": int(self.inference_count), + "last_infer_ts": float(self.last_infer_ts), + "last_timing_ms": dict(self.last_timing_ms), + "trt_cache_dir": str(self.trt_cache_dir), + "contract": self.contract.runtime_summary(), + } + + def _log_runtime_summary(self) -> None: + c = self.contract + level = "DEGRADED" if self.runtime_degraded else "OK" + self.mostrar_log( + "[SEG_RUNTIME] " + f"{level} schema={c.schema} model={self.onnx_path.name} " + f"sha={self.model_sha_short} res={c.width}x{c.height} " + f"input={c.input_name}:{c.input_dtype}:{c.input_layout}:{c.input_color_order} " + f"outputs={self.output_names}" + ) + self.mostrar_log( + "[SEG_RUNTIME] " + f"provider requested={self.provider_requested} primary={self.provider_primary} " + f"chain={self.provider_chain} session={self.session_create_ms:.1f}ms " + f"TRT_cache={self.trt_cache_dir}" + ) + self.mostrar_log( + "[SEG_RUNTIME] " + f"camera_color={self.camera_frame_color} -> model_color={c.input_color_order} | " + f"status={c.status_id2label}" + ) + + def _default_semantic_colormap_rgb(self): + max_id = max(self.contract.semantic_id2label) + colors = [(80, 80, 80)] * (max_id + 1) + # Cores históricas do dataset apenas para preview humano. + colors[self.contract.non_nav_class_id] = (128, 0, 0) + colors[self.contract.nav_class_id] = (0, 128, 0) + return colors + + @staticmethod + def _sha256_file(path: Path, chunk_size: int = 4 * 1024 * 1024) -> str: + h = hashlib.sha256() + with path.open("rb") as f: + while True: + block = f.read(chunk_size) + if not block: + break + h.update(block) + return h.hexdigest()