arquivos restantes do visual worker
This commit is contained in:
parent
257d496496
commit
760d996060
|
|
@ -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,
|
||||||
|
}
|
||||||
|
|
@ -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()
|
||||||
Loading…
Reference in New Issue