diff --git a/AgroBase/AgroBase/Services/MapasService.cs b/AgroBase/AgroBase/Services/MapasService.cs index 6fbd4a14d..9ca907675 100644 --- a/AgroBase/AgroBase/Services/MapasService.cs +++ b/AgroBase/AgroBase/Services/MapasService.cs @@ -745,30 +745,44 @@ namespace AgroBase.Services private static void NormalizarMapa(MapaFeatureCollectionModel dados) { - dados.type = string.IsNullOrWhiteSpace(dados.type) ? "FeatureCollection" : dados.type; + dados.type = string.IsNullOrWhiteSpace(dados.type) + ? "FeatureCollection" + : dados.type; + + var idsUtilizados = new HashSet(StringComparer.Ordinal); for (int i = 0; i < dados.features.Count; i++) { MapaFeatureModel feature = dados.features[i]; - feature.type = string.IsNullOrWhiteSpace(feature.type) ? "Feature" : feature.type; + feature.type = string.IsNullOrWhiteSpace(feature.type) + ? "Feature" + : feature.type; if (feature.properties == null) - { feature.properties = new MapaFeaturePropertiesModel(); - } - /* - * Mantemos a compatibilidade com o fluxo atual: - * a seleção usa IDs internos começando em zero. - */ - string id = (i + 1).ToString(CultureInfo.InvariantCulture); - feature.geometry.id = id; + string id = feature.properties.Id?.Trim(); - if (string.IsNullOrWhiteSpace(feature.properties.Id)) + // Mapa sem ID explícito: + // gera um ID numérico sequencial. + if (string.IsNullOrWhiteSpace(id)) { - feature.properties.Id = id; + id = (i + 1).ToString( + CultureInfo.InvariantCulture + ); } + + if (!idsUtilizados.Add(id)) + { + throw new InvalidOperationException( + $"O mapa possui ID de rua duplicado: '{id}'." + ); + } + + // Uma única identidade em todo o sistema. + feature.properties.Id = id; + feature.geometry.id = id; } } @@ -899,6 +913,15 @@ namespace AgroBase.Services } function obterIdRua(feature, indice) { + if ( + feature && + feature.properties && + feature.properties.Id !== undefined && + feature.properties.Id !== null + ) { + return String(feature.properties.Id); + } + if ( feature && feature.geometry && @@ -916,7 +939,7 @@ namespace AgroBase.Services return String(feature.id); } - return String(indice); + return String(indice + 1); } function ruaSelecionada(id) { diff --git a/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py b/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py index c6c419b6a..5af3b30bf 100644 --- a/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py +++ b/Python/OAK/datasets/oak-fcc-3/_10_export_onnx.py @@ -1,7 +1,7 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- -""" +r""" _10_export_onnx.py Exporta o checkpoint PyTorch do SegFormer Multi-Head OAK-FCC-3 para ONNX. @@ -10,7 +10,7 @@ Exemplo: python .\_10_export_onnx.py --config .\config.json --train-script .\_8_train_multihead_v2.py --checkpoint .\backup\segformer_b1\2026_09_18\stacked_raw5\best_score.pt --out .\backup\segformer_b1\2026_09_18\stacked_raw5\best_score_audit.onnx --opset 17 --device cuda --include-norm --norm-stats .\backup\segformer_b1\2026_09_18\stacked_raw5\norm_stats.json --postprocess none -python .\_10_export_onnx.py --config .\config.json --train-script .\_8_train_multihead_v2.py --checkpoint .\backup\segformer_b1\2026_09_18\stacked_raw5\best_operational.pt --out .\backup\segformer_b1\2026_09_18\stacked_raw5\best_operation.onnx --opset 17 --device cuda --include-norm --norm-stats .\backup\segformer_b1\2026_09_18\stacked_raw5\norm_stats.json --postprocess argmax_fullres +python .\_10_export_onnx.py --config .\config.json --train-script .\_8_train_multihead_v2.py --checkpoint .\backup\segformer_b1\2026_09_18\stacked_raw5\best_operational.pt --out .\backup\segformer_b1\2026_09_18\stacked_raw5\best_operation.onnx --opset 17 --device cuda --include-norm --norm-stats .\backup\segformer_b1\2026_09_18\stacked_raw5\norm_stats.json --postprocess argmax_fullres --target-threshold auto --vegetation-threshold 0.5 --cana-threshold 0.5 """ @@ -178,6 +178,95 @@ def load_selected_norm_stats( return path, mean_sel, std_sel, requested +def _coerce_threshold(value): + """Converte float/dict de metadata em threshold [0, 1].""" + if isinstance(value, dict): + value = value.get("threshold") + try: + if value is None: + return None + v = float(value) + if math.isfinite(v) and 0.0 <= v <= 1.0: + return v + except Exception: + pass + return None + + +def extract_operational_threshold_from_checkpoint(ckpt: dict): + """ + Resolve o threshold operacional da head target salvo pelo trainer. + + Suporta o contrato atual e alguns formatos anteriores usados nos + checkpoints/auditores do projeto. + """ + if not isinstance(ckpt, dict): + return None, "none" + + candidates = [] + + best = ckpt.get("best", {}) or {} + if isinstance(best, dict): + candidates.extend([ + ("best.operational_threshold", best.get("operational_threshold")), + ("best.best_operational_threshold", best.get("best_operational_threshold")), + ]) + + extra = ckpt.get("extra", {}) or {} + if isinstance(extra, dict): + score_now = extra.get("score_now", {}) or {} + if isinstance(score_now, dict): + candidates.extend([ + ("extra.score_now.operational_threshold", score_now.get("operational_threshold")), + ("extra.score_now.best_operational_threshold", score_now.get("best_operational_threshold")), + ]) + + val = extra.get("val", {}) or {} + if isinstance(val, dict): + agri = val.get("agricultural", {}) or {} + if isinstance(agri, dict): + candidates.append(( + "extra.val.agricultural.best_operational_threshold", + agri.get("best_operational_threshold"), + )) + + for source, value in candidates: + threshold = _coerce_threshold(value) + if threshold is not None: + return threshold, source + + return None, "none" + + +def resolve_threshold_arg(raw_value, *, head_name: str, ckpt: dict, allow_auto: bool): + """Resolve CLI numérica ou `auto` e devolve (valor, fonte).""" + raw = str(raw_value).strip().lower() + + if raw == "auto": + if not allow_auto: + raise RuntimeError( + f"--{head_name}-threshold=auto não é suportado para esta head. " + "Informe um valor explícito entre 0 e 1." + ) + value, source = extract_operational_threshold_from_checkpoint(ckpt) + if value is None: + raise RuntimeError( + "Não encontrei o operational_threshold no checkpoint para a head target. " + "Para export de campo não vou cair silenciosamente para 0.5. " + "Informe --target-threshold explicitamente ou use um checkpoint " + "que contenha best.operational_threshold." + ) + return value, f"checkpoint:{source}" + + value = _coerce_threshold(raw) + if value is None: + raise RuntimeError( + f"--{head_name}-threshold inválido: {raw_value!r}. Use valor entre 0 e 1" + + (" ou 'auto'" if allow_auto else "") + "." + ) + return value, "cli" + + # ============================================================ # Wrapper ONNX # ============================================================ @@ -189,8 +278,8 @@ class MultiHeadOnnxWrapper(nn.Module): Pode exportar: - logits crus - logits redimensionados - - argmax em baixa resolução - - argmax em resolução da entrada + - máscara em baixa resolução (threshold explícito nas heads binárias) + - máscara em resolução da entrada (threshold explícito nas heads binárias) Também pode embutir a normalização: x = (x - mean) / std @@ -205,6 +294,7 @@ class MultiHeadOnnxWrapper(nn.Module): norm_mean: List[float] | None = None, norm_std: List[float] | None = None, postprocess: str = "none", + binary_thresholds: Dict[str, float] | None = None, ): super().__init__() self.model = model @@ -212,6 +302,17 @@ class MultiHeadOnnxWrapper(nn.Module): self.resize_to_input = bool(resize_to_input) self.include_norm = bool(include_norm) self.postprocess = str(postprocess).lower() + self.binary_thresholds = { + str(k).strip().lower(): float(v) + for k, v in (binary_thresholds or {}).items() + } + invalid_threshold_heads = sorted( + set(self.binary_thresholds) - {"vegetation", "cana", "target"} + ) + if invalid_threshold_heads: + raise RuntimeError( + f"Threshold binário configurado para heads inválidas: {invalid_threshold_heads}" + ) if self.postprocess not in ("none", "resize_logits", "argmax_lowres", "argmax_fullres"): raise RuntimeError(f"postprocess inválido: {self.postprocess}") @@ -232,6 +333,21 @@ class MultiHeadOnnxWrapper(nn.Module): self.register_buffer("norm_mean", torch.empty(0)) self.register_buffer("norm_std", torch.empty(0)) + def _mask_from_logits(self, head_name: str, logits: torch.Tensor) -> torch.Tensor: + """ + Heads binárias configuradas usam probabilidade explícita da classe 1. + Demais heads preservam argmax tradicional. + """ + threshold = self.binary_thresholds.get(str(head_name).strip().lower()) + if threshold is not None: + # vegetation/cana/target são heads binárias de 2 logits pelo contrato + # do modelo. Evitamos branch dependente de shape aqui para manter o + # trace/export ONNX totalmente estático. + prob_positive = torch.softmax(logits, dim=1)[:, 1, :, :] + return (prob_positive >= threshold).to(torch.uint8) + + return torch.argmax(logits, dim=1).to(torch.uint8) + def forward(self, pixel_values: torch.Tensor): input_hw = pixel_values.shape[-2:] @@ -262,7 +378,7 @@ class MultiHeadOnnxWrapper(nn.Module): result.append(logits) elif self.postprocess == "argmax_lowres": - mask = torch.argmax(logits, dim=1).to(torch.uint8) + mask = self._mask_from_logits(head_name, logits) result.append(mask) elif self.postprocess == "argmax_fullres": @@ -272,7 +388,7 @@ class MultiHeadOnnxWrapper(nn.Module): mode="bilinear", align_corners=False, ) - mask = torch.argmax(logits, dim=1).to(torch.uint8) + mask = self._mask_from_logits(head_name, logits) result.append(mask) return tuple(result) @@ -331,6 +447,32 @@ def main(): ), ) + parser.add_argument( + "--target-threshold", + default="auto", + help=( + "Threshold da probabilidade da classe positiva da head target quando o " + "postprocess gera máscara. Padrão='auto': lê o operational_threshold " + "salvo no checkpoint e falha se ele não existir." + ), + ) + parser.add_argument( + "--vegetation-threshold", + default="0.5", + help=( + "Threshold da classe vegetation para export mask. Padrão=0.5. " + "Use override explícito apenas após calibração específica dessa head." + ), + ) + parser.add_argument( + "--cana-threshold", + default="0.5", + help=( + "Threshold da classe cana para export mask. Padrão=0.5. " + "Use override explícito apenas após calibração específica dessa head." + ), + ) + parser.add_argument( "--dynamic-batch", action="store_true", @@ -462,6 +604,58 @@ def main(): model.to(device) model.eval() + mask_postprocess = args.postprocess in ("argmax_lowres", "argmax_fullres") + checkpoint_operational_threshold, checkpoint_operational_threshold_source = ( + extract_operational_threshold_from_checkpoint(ckpt) + ) + + binary_thresholds = {} + threshold_sources = {} + + if mask_postprocess: + target_threshold, target_threshold_source = resolve_threshold_arg( + args.target_threshold, + head_name="target", + ckpt=ckpt, + allow_auto=True, + ) + vegetation_threshold, vegetation_threshold_source = resolve_threshold_arg( + args.vegetation_threshold, + head_name="vegetation", + ckpt=ckpt, + allow_auto=False, + ) + cana_threshold, cana_threshold_source = resolve_threshold_arg( + args.cana_threshold, + head_name="cana", + ckpt=ckpt, + allow_auto=False, + ) + + binary_thresholds = { + "target": target_threshold, + "vegetation": vegetation_threshold, + "cana": cana_threshold, + } + threshold_sources = { + "target": target_threshold_source, + "vegetation": vegetation_threshold_source, + "cana": cana_threshold_source, + } + + print("[THRESHOLD] Máscaras binárias serão exportadas com thresholds explícitos:") + for head_name in ("target", "vegetation", "cana"): + print( + f" {head_name:10s}: {binary_thresholds[head_name]:.3f} " + f"({threshold_sources[head_name]})" + ) + else: + print( + "[THRESHOLD] postprocess retorna logits; thresholds não são aplicados no grafo. " + f"Checkpoint operational={checkpoint_operational_threshold} " + f"source={checkpoint_operational_threshold_source}" + ) + norm_stats_path = None norm_mean = None norm_std = None @@ -488,6 +682,7 @@ def main(): norm_mean=norm_mean, norm_std=norm_std, postprocess=args.postprocess, + binary_thresholds=binary_thresholds, ) wrapper.to(device) wrapper.eval() @@ -564,6 +759,11 @@ def main(): "heads_config": heads_config, "checkpoint_epoch": ckpt.get("epoch", None), "checkpoint_best": ckpt.get("best", None), + "checkpoint_operational_threshold": checkpoint_operational_threshold, + "checkpoint_operational_threshold_source": checkpoint_operational_threshold_source, + "thresholds_applied_in_graph": bool(mask_postprocess), + "binary_thresholds": binary_thresholds, + "binary_threshold_sources": threshold_sources, "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, @@ -588,6 +788,10 @@ def main(): "oak.output_heads": json.dumps(output_heads), "oak.output_names": json.dumps(output_names), "oak.stats_source_tag": experiment_tag(config, channels), + "oak.thresholds_applied_in_graph": json.dumps(bool(mask_postprocess)), + "oak.binary_thresholds": json.dumps(binary_thresholds), + "oak.binary_threshold_sources": json.dumps(threshold_sources), + "oak.checkpoint_operational_threshold": json.dumps(checkpoint_operational_threshold), } existing = {item.key: item for item in onnx_model.metadata_props} for key, value in metadata_values.items(): diff --git a/Python/OAK/datasets/oak-fcc-3/_1_weeds_pair_sorter.py b/Python/OAK/datasets/oak-fcc-3/_1_weeds_pair_sorter.py index 4b376d14a..573d4d9df 100644 --- a/Python/OAK/datasets/oak-fcc-3/_1_weeds_pair_sorter.py +++ b/Python/OAK/datasets/oak-fcc-3/_1_weeds_pair_sorter.py @@ -20,6 +20,12 @@ Fluxo: - Ao classificar, salva o PNG derivado em previews/ e registra o contrato no meta. - Opcionalmente mostra CAM_A/CAM_B/CAM_C nativos direto dos .bin. - Você usa teclas 1..9/0 para enviar o conjunto da amostra para uma label. +- Opcionalmente recebe um modelo .onnx ou checkpoint .pt e mostra a predição do + modelo. No modo full, a prediction é sobreposta ao RGB com transparência ajustável + de 0% (RGB puro) a 100% (prediction pura), sempre sem cortar o frame. +- Quando há modelo, salva predictions/.png no MESMO domínio RGB canônico + da máscara humana. Fora do crop inferido pelo modelo os pixels ficam brancos + (ignore), nunca são inventados como chão. Estrutura de saída: out_root/ @@ -28,12 +34,14 @@ Estrutura de saída: previews/ metas/ masks/ + predictions/ # criada somente quando um modelo é informado Regras: - PNG/JPG vai para previews/ - JSON vai para metas/ - BINs vão para bins/ - masks/ é criada vazia +- predictions/ recebe a máscara prevista no domínio RGB canônico, com ignore branco fora do crop Dependências: pip install pillow numpy @@ -46,13 +54,16 @@ import csv import json import math import hashlib +import importlib.util +import os import re import shutil import sys +import time from dataclasses import dataclass, field from datetime import datetime from pathlib import Path -from typing import Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Sequence, Tuple import tkinter as tk from tkinter import filedialog, messagebox @@ -94,6 +105,31 @@ OLD_CAM_TO_CAMERA = { } CAM_ORDER = {"CAM_A": 0, "CAM_B": 1, "CAM_C": 2} REQUIRED_CAMERAS = ("CAM_A", "CAM_B", "CAM_C") +RAW5_CHANNEL_ORDER = ("R", "G", "B", "RE", "NIR") +PREDICTION_IGNORE_ID = 255 +SEMANTIC_PREDICTION_COLORS_RGB = { + 0: (128, 0, 0), # chao + 1: (0, 0, 128), # cana + 2: (0, 128, 0), # erva + PREDICTION_IGNORE_ID: (255, 255, 255), +} +BINARY_PREDICTION_COLORS_RGB = { + "vegetation": { + 0: (25, 25, 25), + 1: (0, 180, 0), + PREDICTION_IGNORE_ID: (255, 255, 255), + }, + "cana": { + 0: (25, 25, 25), + 1: (0, 0, 180), + PREDICTION_IGNORE_ID: (255, 255, 255), + }, + "target": { + 0: (25, 25, 25), + 1: (0, 220, 0), + PREDICTION_IGNORE_ID: (255, 255, 255), + }, +} GENERATED_PREVIEW_VERSION = "weed_pair_sorter_canonical_preview_v4_2026_09_14_bayer_explicit" ANNOTATION_PREVIEW_CONFIG = { "schema": "rgb_raw_annotation_preview_v2", @@ -133,6 +169,12 @@ class SampleBundle: bin_by_camera: Dict[str, Path] = field(default_factory=dict) +@dataclass +class PredictionArtifact: + image_rgb: np.ndarray + info: dict + + # ============================================================ # Utilidades de arquivo/meta # ============================================================ @@ -799,6 +841,7 @@ def choose_unique_output_base(label_root: Path, desired_base: str) -> str: candidates = [ label_root / "previews" / f"{base}.png", label_root / "metas" / f"{base}.json", + label_root / "predictions" / f"{base}.png", *[label_root / "bins" / f"{base}_{cam}.bin" for cam in REQUIRED_CAMERAS], ] if not any(c.exists() for c in candidates): @@ -806,7 +849,15 @@ def choose_unique_output_base(label_root: Path, desired_base: str) -> str: k += 1 -def write_destination_meta(src_meta_path: Path, dst_meta_path: Path, dst_preview_name: str, preview_report: dict, preview_origin: str, payload_files: Optional[Dict[str, str]] = None): +def write_destination_meta( + src_meta_path: Path, + dst_meta_path: Path, + dst_preview_name: str, + preview_report: dict, + preview_origin: str, + payload_files: Optional[Dict[str, str]] = None, + prediction_info: Optional[dict] = None, +): meta = load_json(src_meta_path) meta["saved_preview_path"] = dst_preview_name meta["saved_preview_method"] = "raw_processor_preview_rgb_beauty_canonical_sorter" @@ -817,6 +868,8 @@ def write_destination_meta(src_meta_path: Path, dst_meta_path: Path, dst_preview meta["annotation_preview"] = preview_report meta["sorting_preview_origin"] = preview_origin meta["sorting_preview_version"] = GENERATED_PREVIEW_VERSION + if prediction_info: + meta["model_prediction"] = dict(prediction_info) tmp = dst_meta_path.with_suffix(dst_meta_path.suffix + ".tmp") tmp.parent.mkdir(parents=True, exist_ok=True) with tmp.open("w", encoding="utf-8") as f: @@ -1018,6 +1071,918 @@ def build_bin_preview_image(bundle: SampleBundle, cam_id: str, module_params_ove raise RuntimeError(f"formato não suportado para {cam_id}: dtype={arr.dtype} shape={arr.shape}") + +# ============================================================ +# Predição opcional do modelo (.onnx ou .pt) +# ============================================================ + +def sha256_file(path: Path, chunk_size: int = 1024 * 1024) -> str: + h = hashlib.sha256() + with path.open("rb") as f: + while True: + chunk = f.read(chunk_size) + if not chunk: + break + h.update(chunk) + return h.hexdigest() + + +def _import_raw_processor_core(): + """Importa a implementação oficial usada no normalize/runtime.""" + try: + from core.raw_processor_core import RawProcessorCore + return RawProcessorCore + except Exception as first_exc: + here = Path(__file__).resolve() + for parent in [here.parent, *here.parents]: + candidate = parent / "core" / "raw_processor_core.py" + if candidate.is_file(): + if str(parent) not in sys.path: + sys.path.insert(0, str(parent)) + try: + from core.raw_processor_core import RawProcessorCore + return RawProcessorCore + except Exception: + pass + raise RuntimeError( + "Não consegui importar core.raw_processor_core.RawProcessorCore. " + "Execute o sorter dentro do projeto oak-fcc-3 ou ajuste o PYTHONPATH. " + f"Erro original: {first_exc}" + ) from first_exc + + +def _import_python_module(path: Path, module_name: str): + if not path.is_file(): + raise FileNotFoundError(f"Script Python não encontrado: {path}") + spec = importlib.util.spec_from_file_location(module_name, str(path.resolve())) + if spec is None or spec.loader is None: + raise RuntimeError(f"Não consegui importar módulo Python: {path}") + mod = importlib.util.module_from_spec(spec) + sys.modules[module_name] = mod + spec.loader.exec_module(mod) + return mod + + +def _resolve_local_path(value: Optional[Path | str], base: Optional[Path] = None) -> Optional[Path]: + if value is None or str(value).strip() == "": + return None + p = Path(value).expanduser() + if p.is_absolute(): + return p.resolve() + candidates = [] + if base is not None: + candidates.append(Path(base) / p) + candidates.extend([Path.cwd() / p, Path(__file__).resolve().parent / p]) + for c in candidates: + if c.exists(): + return c.resolve() + return candidates[0].resolve() if candidates else p.resolve() + + +def _extract_operational_threshold_from_dict(data: dict) -> Tuple[Optional[float], str]: + """Mesma ordem usada pelos scripts de review/export atuais.""" + if not isinstance(data, dict): + return None, "none" + + direct = [ + ("operational_threshold", data.get("operational_threshold")), + ("checkpoint_operational_threshold", data.get("checkpoint_operational_threshold")), + ] + best = data.get("checkpoint_best", data.get("best", {})) or {} + if isinstance(best, dict): + direct.extend([ + ("best.operational_threshold", best.get("operational_threshold")), + ("best.best_operational_threshold", best.get("best_operational_threshold")), + ]) + + extra = data.get("extra", {}) or {} + if isinstance(extra, dict): + score_now = extra.get("score_now", {}) or {} + if isinstance(score_now, dict): + direct.append(("extra.score_now.operational_threshold", score_now.get("operational_threshold"))) + val = extra.get("val", {}) or {} + if isinstance(val, dict): + agri = val.get("agricultural", {}) or {} + if isinstance(agri, dict): + bop = agri.get("best_operational_threshold", {}) or {} + if isinstance(bop, dict): + direct.append(("extra.val.agricultural.best_operational_threshold.threshold", bop.get("threshold"))) + + for source, value in direct: + try: + if value is not None: + v = float(value) + if 0.0 <= v <= 1.0: + return v, source + except Exception: + pass + return None, "none" + + +def _load_norm_stats_simple(path: Path, channel_names: Sequence[str]) -> Tuple[List[float], List[float]]: + js = load_json(path) + mean = js.get("mean") + std = js.get("std") + names = [str(x).strip().upper() for x in js.get("channels", [])] + requested = [str(x).strip().upper() for x in channel_names] + + if mean is None or std is None: + raise RuntimeError(f"norm_stats inválido, faltando mean/std: {path}") + if names: + missing = [x for x in requested if x not in names] + if missing: + raise RuntimeError(f"norm_stats não possui canais {missing}: {path}") + idx = [names.index(x) for x in requested] + elif len(mean) == len(requested): + idx = list(range(len(requested))) + elif len(mean) == len(RAW5_CHANNEL_ORDER): + idx = [RAW5_CHANNEL_ORDER.index(x) for x in requested] + else: + raise RuntimeError( + f"Não é seguro mapear norm_stats: mean={len(mean)} canais pedidos={requested}" + ) + + mean_sel = [float(mean[i]) for i in idx] + std_sel = [float(std[i]) for i in idx] + if any((not np.isfinite(x)) for x in mean_sel + std_sel) or any(x <= 0 for x in std_sel): + raise RuntimeError(f"norm_stats contém valores inválidos: {path}") + return mean_sel, std_sel + + +def _softmax_np(logits: np.ndarray, axis: int = 0) -> np.ndarray: + x = np.asarray(logits, dtype=np.float32) + x = x - np.max(x, axis=axis, keepdims=True) + ex = np.exp(x) + return ex / np.maximum(np.sum(ex, axis=axis, keepdims=True), 1e-12) + + +def _resize_logits_chw(logits_chw: np.ndarray, hw: Tuple[int, int]) -> np.ndarray: + h, w = map(int, hw) + if logits_chw.shape[-2:] == (h, w): + return logits_chw.astype(np.float32, copy=False) + return np.stack( + [ + cv2.resize( + logits_chw[c].astype(np.float32), + (w, h), + interpolation=cv2.INTER_LINEAR, + ) + for c in range(logits_chw.shape[0]) + ], + axis=0, + ) + + +def prediction_ids_to_rgb_domain( + pred_ids: np.ndarray, + fusion_result: Optional[dict], + rgb_size: Tuple[int, int], + head: str, + ignore_id: int = PREDICTION_IGNORE_ID, +) -> Tuple[np.ndarray, dict]: + """ + Volta a prediction do espaço final do modelo para o RGB canônico de anotação. + + Contrato: + RGB canônico -> ref_shape -> crop_box -> resize -> tensor/modelo + prediction -> resize NEAREST inverso -> crop_box no ref_shape -> RGB canônico + + Fora do crop a saída é IGNORE branco, nunca chão artificial. + """ + if cv2 is None: + raise RuntimeError("OpenCV é obrigatório para reprojetar a prediction.") + + pred = np.asarray(pred_ids, dtype=np.uint8) + if pred.ndim != 2: + raise RuntimeError(f"Prediction 2D esperada, recebido {pred.shape}") + + rgb_w, rgb_h = map(int, rgb_size) + if rgb_w <= 0 or rgb_h <= 0: + raise RuntimeError(f"Tamanho RGB inválido: {rgb_size}") + + fr = fusion_result if isinstance(fusion_result, dict) else {} + ref_shape = fr.get("ref_shape") + if isinstance(ref_shape, (list, tuple)) and len(ref_shape) == 2: + ref_h, ref_w = int(ref_shape[0]), int(ref_shape[1]) + else: + ref_h, ref_w = rgb_h, rgb_w + + if ref_w <= 0 or ref_h <= 0: + ref_w, ref_h = rgb_w, rgb_h + + crop_box = fr.get("crop_box") + if isinstance(crop_box, (list, tuple)) and len(crop_box) == 4: + x0, y0, x1, y1 = [int(v) for v in crop_box] + else: + x0, y0, x1, y1 = 0, 0, ref_w, ref_h + + x0 = max(0, min(ref_w - 1, x0)) + y0 = max(0, min(ref_h - 1, y0)) + x1 = max(x0 + 1, min(ref_w, x1)) + y1 = max(y0 + 1, min(ref_h, y1)) + + crop_w = x1 - x0 + crop_h = y1 - y0 + pred_crop = cv2.resize(pred, (crop_w, crop_h), interpolation=cv2.INTER_NEAREST) + + ids_ref = np.full((ref_h, ref_w), int(ignore_id), dtype=np.uint8) + ids_ref[y0:y1, x0:x1] = pred_crop[:crop_h, :crop_w] + + if (ref_w, ref_h) != (rgb_w, rgb_h): + ids_rgb = cv2.resize(ids_ref, (rgb_w, rgb_h), interpolation=cv2.INTER_NEAREST) + else: + ids_rgb = ids_ref + + head = str(head).strip().lower() + if head == "semantic": + palette = SEMANTIC_PREDICTION_COLORS_RGB + else: + palette = BINARY_PREDICTION_COLORS_RGB.get(head) + if palette is None: + raise RuntimeError(f"Head de prediction não suportada: {head}") + + out = np.full((rgb_h, rgb_w, 3), 255, dtype=np.uint8) + for cls_id, color in palette.items(): + out[ids_rgb == int(cls_id)] = np.array(color, dtype=np.uint8) + + geom = { + "coordinate_space": "rgb_canonical_annotation_space", + "ignore_id": int(ignore_id), + "ignore_color_rgb": [255, 255, 255], + "outside_prediction_policy": "white_ignore", + "ref_shape_hw": [int(ref_h), int(ref_w)], + "rgb_size_wh": [int(rgb_w), int(rgb_h)], + "crop_box_ref_xyxy": [int(x0), int(y0), int(x1), int(y1)], + "model_prediction_hw": [int(pred.shape[0]), int(pred.shape[1])], + "inverse_interpolation": "nearest", + } + return np.ascontiguousarray(out), geom + + +class PredictionEngine: + """ + Inferência opcional para auxiliar a triagem manual. + + ONNX: + - aceita audit/logits (postprocess none/resize_logits); + - aceita produção mask-only (semantic_mask/target_mask/etc.); + - detecta normalização embutida pelo .export_meta.json. + + PT: + - reutiliza _9_test_multihead_v2.py para reconstruir exatamente o modelo. + """ + + def __init__( + self, + model_path: Path, + config_path: Optional[Path] = None, + test_script_path: Optional[Path] = None, + norm_stats_path: Optional[Path] = None, + provider: str = "tensorrt", + head: str = "semantic", + target_threshold: str = "auto", + ): + if cv2 is None: + raise RuntimeError("OpenCV é obrigatório para usar prediction no sorter.") + + self.model_path = Path(model_path).resolve() + if not self.model_path.is_file(): + raise FileNotFoundError(f"Modelo não encontrado: {self.model_path}") + + self.model_sha256 = sha256_file(self.model_path) + self.config_path = _resolve_local_path(config_path) if config_path else None + self.test_script_path = _resolve_local_path(test_script_path) if test_script_path else None + self.norm_stats_path = _resolve_local_path(norm_stats_path) if norm_stats_path else None + self.provider_requested = str(provider or "tensorrt").strip().lower() + self.head = str(head or "semantic").strip().lower() + if self.head not in ("semantic", "vegetation", "cana", "target"): + raise RuntimeError(f"prediction head inválida: {self.head}") + + self.target_threshold_arg = str(target_threshold or "auto").strip().lower() + self.target_threshold = 0.5 + self.target_threshold_source = "default_0.5" + self.input_channel_names = list(RAW5_CHANNEL_ORDER) + self.model_w = 960 + self.model_h = 600 + self.normalization_embedded = False + self.mean = None + self.std = None + self.backend = self.model_path.suffix.lower().lstrip(".") + self.active_provider = "" + self.output_kind = "" + self.selected_output_name = "" + self._core_cache: Dict[Tuple[int, int, str, str], Any] = {} + self._last_cache_key = None + self._last_artifact: Optional[PredictionArtifact] = None + self._input_diag_printed = False + + if self.model_path.suffix.lower() == ".onnx": + self._init_onnx() + elif self.model_path.suffix.lower() in (".pt", ".pth"): + self._init_pt() + else: + raise RuntimeError("Modelo opcional deve ser .onnx, .pt ou .pth") + + def _resolve_target_threshold(self, metadata: Optional[dict] = None, ckpt_value: Optional[float] = None): + if self.target_threshold_arg != "auto": + value = float(self.target_threshold_arg) + if not 0.0 <= value <= 1.0: + raise RuntimeError("target_threshold deve ficar entre 0 e 1") + self.target_threshold = value + self.target_threshold_source = "cli" + return + + if ckpt_value is not None: + value = float(ckpt_value) + if 0.0 <= value <= 1.0: + self.target_threshold = value + self.target_threshold_source = "checkpoint" + return + + if metadata: + value, source = _extract_operational_threshold_from_dict(metadata) + if value is not None: + self.target_threshold = value + self.target_threshold_source = source + return + + self.target_threshold = 0.5 + self.target_threshold_source = "fallback_0.5" + + def _resolve_norm_stats_auto(self, config: Optional[dict], config_dir: Optional[Path]) -> Optional[Path]: + if self.norm_stats_path is not None: + if not self.norm_stats_path.is_file(): + raise FileNotFoundError(f"norm_stats não encontrado: {self.norm_stats_path}") + return self.norm_stats_path + + candidates: List[Path] = [] + if config is not None and config_dir is not None: + res = config.get("resolucao", [self.model_w, self.model_h]) + w, h = int(res[0]), int(res[1]) + model_name = str(config.get("model_name", "test_multi")) + modelo = str(config.get("modelo", "segformer_b1")) + stats_tag = str(config.get("stats_source_tag", "")).strip() + if not stats_tag: + stats_tag = f"{str(config.get('fusion_mode', 'stacked')).strip()}_raw{len(self.input_channel_names)}" + candidates.extend([ + self.model_path.parent / "norm_stats.json", + config_dir / "dataset" / f"{w}x{h}" / "group" / "norm_stats.json", + config_dir / "backup" / modelo / model_name / stats_tag / "norm_stats.json", + ]) + else: + candidates.append(self.model_path.parent / "norm_stats.json") + + return next((p.resolve() for p in candidates if p.is_file()), None) + + @staticmethod + def _decode_onnx_meta_value(value): + """Decodifica metadata custom do ONNX, preservando string quando não for JSON.""" + if value is None: + return None + text = str(value) + try: + return json.loads(text) + except Exception: + return text + + def _read_onnx_graph_contract(self) -> dict: + """ + Lê o contrato que o exportador grava DENTRO do próprio ONNX. + + O grafo é a autoridade primária porque o arquivo pode ser renomeado/copied + sem o .export_meta.json lateral. O sidecar continua sendo usado para campos + adicionais (mean/std, checkpoint_best etc.) quando existir. + """ + try: + mm = self.session.get_modelmeta() + raw = dict(getattr(mm, "custom_metadata_map", {}) or {}) + except Exception: + raw = {} + + out = {"_onnx_custom_metadata_raw": raw} + mapping = { + "oak.input_channel_names": "input_channel_names", + "oak.normalization_embedded": "include_norm", + "oak.postprocess": "postprocess", + "oak.output_heads": "heads", + "oak.output_names": "output_names", + "oak.stats_source_tag": "stats_source_tag", + "oak.thresholds_applied_in_graph": "thresholds_applied_in_graph", + "oak.binary_thresholds": "binary_thresholds", + "oak.binary_threshold_sources": "binary_threshold_sources", + "oak.checkpoint_operational_threshold": "checkpoint_operational_threshold", + } + for src, dst in mapping.items(): + if src in raw: + out[dst] = self._decode_onnx_meta_value(raw[src]) + + if "include_norm" in out: + out["input_contract"] = { + "normalization_embedded": bool(out["include_norm"]), + } + return out + + def _init_onnx(self): + try: + import onnxruntime as ort + except Exception as exc: + raise RuntimeError( + "onnxruntime não está disponível. Instale/ative onnxruntime-gpu no venv." + ) from exc + + # Sidecar é útil, mas NÃO pode ser a única fonte de contrato: em produção + # é comum renomear best_operation.onnx -> model-4_0.onnx sem renomear o JSON. + meta_path = self.model_path.with_suffix(".export_meta.json") + sidecar_meta = load_json(meta_path) if meta_path.is_file() else {} + self.export_meta_path = meta_path if meta_path.is_file() else None + + available = set(ort.get_available_providers()) + providers: List[Any] = [] + req = self.provider_requested + + if req in ("tensorrt", "trt") and "TensorrtExecutionProvider" in available: + cache_dir = ( + self.model_path.parent + / "trt_cache_sorter" + / f"{self.model_path.stem}_{self.model_sha256[:16]}" + ) + cache_dir.mkdir(parents=True, exist_ok=True) + providers.append(( + "TensorrtExecutionProvider", + { + "device_id": 0, + "trt_fp16_enable": True, + "trt_engine_cache_enable": True, + "trt_engine_cache_path": str(cache_dir), + "trt_timing_cache_enable": True, + "trt_timing_cache_path": str(cache_dir), + }, + )) + + if req in ("tensorrt", "trt", "cuda", "gpu") and "CUDAExecutionProvider" in available: + providers.append("CUDAExecutionProvider") + if "CPUExecutionProvider" in available: + providers.append("CPUExecutionProvider") + + if not providers: + raise RuntimeError( + f"Nenhum provider ONNX utilizável. solicitado={req} disponíveis={sorted(available)}" + ) + + opts = ort.SessionOptions() + opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL + self.session = ort.InferenceSession( + str(self.model_path), + sess_options=opts, + providers=providers, + ) + self.active_provider = ",".join(self.session.get_providers()) + self.input_name = self.session.get_inputs()[0].name + + # O contrato interno do ONNX prevalece sobre o sidecar para tudo que + # descreve o próprio grafo. O sidecar fornece informações extras. + graph_meta = self._read_onnx_graph_contract() + self.export_meta = dict(sidecar_meta) + for k, v in graph_meta.items(): + if k != "_onnx_custom_metadata_raw" and v is not None: + if k == "input_contract": + merged = dict(self.export_meta.get("input_contract", {}) or {}) + merged.update(v) + self.export_meta["input_contract"] = merged + else: + self.export_meta[k] = v + self.onnx_graph_meta = graph_meta + + exported_names = [ + str(x).strip().upper() + for x in self.export_meta.get("input_channel_names", []) + ] + if exported_names: + self.input_channel_names = exported_names + + input_shape = self.session.get_inputs()[0].shape + if len(input_shape) != 4: + raise RuntimeError(f"Entrada ONNX inesperada: {input_shape}") + if isinstance(input_shape[-1], int) and isinstance(input_shape[-2], int): + self.model_h = int(input_shape[-2]) + self.model_w = int(input_shape[-1]) + elif self.export_meta.get("input_shape"): + shape = self.export_meta["input_shape"] + self.model_h = int(shape[-2]) + self.model_w = int(shape[-1]) + + input_contract = self.export_meta.get("input_contract", {}) or {} + self.normalization_embedded = bool( + input_contract.get( + "normalization_embedded", + self.export_meta.get("include_norm", False), + ) + ) + + # Só normaliza externamente quando o próprio ONNX declara que precisa. + norm_mean = self.export_meta.get("norm_mean") + norm_std = self.export_meta.get("norm_std") + if not self.normalization_embedded: + if norm_mean is not None and norm_std is not None: + self.mean = np.asarray(norm_mean, dtype=np.float32).reshape(len(self.input_channel_names), 1, 1) + self.std = np.asarray(norm_std, dtype=np.float32).reshape(len(self.input_channel_names), 1, 1) + else: + cfg = load_json(self.config_path) if self.config_path and self.config_path.is_file() else None + cfg_dir = self.config_path.parent if self.config_path and self.config_path.is_file() else None + npath = self._resolve_norm_stats_auto(cfg, cfg_dir) + if npath is None: + raise RuntimeError( + "ONNX declara normalização EXTERNA e não encontrei norm_stats. " + "Informe --prediction-norm-stats." + ) + mean, std = _load_norm_stats_simple(npath, self.input_channel_names) + self.mean = np.asarray(mean, dtype=np.float32).reshape(len(mean), 1, 1) + self.std = np.asarray(std, dtype=np.float32).reshape(len(std), 1, 1) + self.norm_stats_path = npath + + output_names = [o.name for o in self.session.get_outputs()] + exact_mask = f"{self.head}_mask" + exact_logits = f"{self.head}_logits" + + if exact_mask in output_names: + self.selected_output_name = exact_mask + self.output_kind = "mask" + elif exact_logits in output_names: + self.selected_output_name = exact_logits + self.output_kind = "logits" + else: + candidates = [x for x in output_names if self.head in x.lower()] + if not candidates: + raise RuntimeError( + f"Head {self.head!r} não encontrada no ONNX. outputs={output_names}" + ) + self.selected_output_name = candidates[0] + self.output_kind = "mask" if "mask" in candidates[0].lower() else "logits" + + if self.head == "target": + if self.output_kind == "mask": + bt = self.export_meta.get("binary_thresholds", {}) or {} + value = bt.get("target") if isinstance(bt, dict) else None + if value is not None: + self.target_threshold = float(value) + self.target_threshold_source = "onnx_embedded" + else: + self._resolve_target_threshold(self.export_meta) + else: + self._resolve_target_threshold(self.export_meta) + + contract_source = "onnx_graph" + if sidecar_meta: + contract_source += "+sidecar" + print( + f"[SORTER_MODEL] ONNX={self.model_path.name} head={self.head} " + f"output={self.selected_output_name} kind={self.output_kind} " + f"provider={self.active_provider} input={self.model_w}x{self.model_h}" + ) + print( + f"[SORTER_MODEL][CONTRACT] source={contract_source} " + f"norm_embedded={self.normalization_embedded} " + f"channels={self.input_channel_names} " + f"postprocess={self.export_meta.get('postprocess', '?')} " + f"sidecar={'sim' if sidecar_meta else 'não'}" + ) + def _init_pt(self): + if self.config_path is None or not self.config_path.is_file(): + raise FileNotFoundError( + "Checkpoint .pt exige config.json. Informe --prediction-config." + ) + if self.test_script_path is None: + self.test_script_path = _resolve_local_path("_9_test_multihead_v2.py", self.config_path.parent) + if self.test_script_path is None or not self.test_script_path.is_file(): + raise FileNotFoundError( + "Checkpoint .pt exige _9_test_multihead_v2.py. Informe --prediction-test-script." + ) + + try: + import torch + except Exception as exc: + raise RuntimeError("PyTorch não está disponível para abrir checkpoint .pt") from exc + + base = _import_python_module( + self.test_script_path, + f"weed_pair_sorter_test_{self.model_sha256[:10]}", + ) + self._pt_base = base + config = load_json(self.config_path) + config["runtime_mode"] = "all" + self.input_channel_names = [ + str(x).strip().upper() + for x in base.get_input_channel_names(config) + ] + + res = config.get("resolucao", [960, 600]) + self.model_w, self.model_h = int(res[0]), int(res[1]) + + labelmap = self.config_path.parent / "dataset" / "labelmap.txt" + semantic_id2label, semantic_label2id, ignore_id, _ = base.load_labelmap(labelmap) + heads_config = base.build_heads_config(config, ignore_index=int(ignore_id)) + heads_config["semantic"]["num_classes"] = int(len(semantic_id2label)) + heads_config["semantic"]["ignore_index"] = int(ignore_id) + + npath = self._resolve_norm_stats_auto(config, self.config_path.parent) + if npath is None: + raise FileNotFoundError( + "Não encontrei norm_stats para o checkpoint .pt. " + "Informe --prediction-norm-stats." + ) + mean, std = base.load_norm_stats(npath, channel_names=self.input_channel_names) + if mean is None or std is None: + raise RuntimeError(f"Falha ao carregar norm_stats: {npath}") + self.norm_stats_path = npath + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + self._pt_tester = base.MultiHeadTester( + config=config, + ckpt_path=self.model_path, + device=device, + input_channel_names=self.input_channel_names, + heads_config=heads_config, + semantic_id2label=semantic_id2label, + semantic_label2id=semantic_label2id, + mean=mean, + std=std, + use_amp=True, + ) + self.active_provider = str(device) + self.output_kind = "runtime_prediction" + self.selected_output_name = self.head + + if self.head == "target": + self._resolve_target_threshold( + ckpt_value=getattr(self._pt_tester, "ckpt_operational_threshold", None) + ) + self._pt_tester.target_threshold = float(self.target_threshold) + + print( + f"[SORTER_MODEL] PT={self.model_path.name} head={self.head} " + f"device={device} input={self.model_w}x{self.model_h}" + ) + + def _build_tensor( + self, + bundle: SampleBundle, + module_params_override: Optional[Path], + ) -> Tuple[np.ndarray, Optional[dict], Path]: + meta = load_json(bundle.meta_path) + mp_path = resolve_module_params_path( + bundle.meta_path, + meta, + module_params_override, + ) + module_params = load_json(mp_path) + bayer, _ = _authoritative_rgb_bayer(meta, module_params) + + frame: Dict[str, np.ndarray] = {} + for cam in REQUIRED_CAMERAS: + if cam not in bundle.bin_by_camera: + raise RuntimeError(f"{cam} ausente para inferência") + frame[cam] = load_camera_payload(meta, bundle.bin_by_camera[cam], cam) + + rgb_w, rgb_h = _rgb_native_size(meta) + RawProcessorCore = _import_raw_processor_core() + key = (int(rgb_w), int(rgb_h), str(bayer), str(mp_path.resolve())) + if key not in self._core_cache: + self._core_cache[key] = RawProcessorCore( + sensor_width=int(rgb_w), + sensor_height=int(rgb_h), + bayer_pattern=str(bayer), + calibration_json_path=str(mp_path), + ) + core = self._core_cache[key] + + stream_meta = meta.get("stream_meta", {}) or {} + processing_meta = dict(stream_meta) + processing_meta["frame_type"] = "RAW_BRUTO" + if "camera_info" not in processing_meta and isinstance(meta.get("camera_info"), dict): + processing_meta["camera_info"] = meta.get("camera_info") + for key_name in ( + "actual_camera_controls", + "startup_camera_controls", + "radiometric_last_result", + ): + if meta.get(key_name) is not None: + processing_meta[key_name] = meta.get(key_name) + + # Replica o normalize/worker: deixa o próprio RawProcessorCore fazer o + # resize final para a resolução do modelo. Isso preserva INTER_AREA no + # downsample e evita uma implementação geométrica paralela no sorter. + try: + tensor5 = core.build_infer_tensor_from_stream( + frame, processing_meta, 5, target_size=(self.model_w, self.model_h) + ) + except TypeError: + # Compatibilidade com core legado sem target_size na assinatura. + tensor5 = core.build_infer_tensor_from_stream(frame, processing_meta, 5) + if tensor5 is None: + raise RuntimeError(f"RawProcessorCore retornou None para {bundle.sample_id}") + + tensor5 = np.asarray(tensor5, dtype=np.float32) + if tensor5.ndim != 3: + raise RuntimeError(f"Tensor inferência inválido: {tensor5.shape}") + if tensor5.shape[0] != 5 and tensor5.shape[-1] == 5: + tensor5 = np.transpose(tensor5, (2, 0, 1)) + if tensor5.shape[0] < 5: + raise RuntimeError(f"Tensor Raw5 esperado, recebido {tensor5.shape}") + + unsupported = [x for x in self.input_channel_names if x not in RAW5_CHANNEL_ORDER] + if unsupported: + raise RuntimeError( + "Sorter prediction atualmente aceita canais físicos Raw5. " + f"Modelo pede canais derivados/não suportados={unsupported}" + ) + idx = [RAW5_CHANNEL_ORDER.index(x) for x in self.input_channel_names] + tensor = np.ascontiguousarray(tensor5[idx], dtype=np.float32) + + if tensor.shape[-2:] != (self.model_h, self.model_w): + # Fallback apenas para RawProcessorCore legado. Usa a mesma política + # do core oficial: INTER_AREA para reduzir, INTER_LINEAR para ampliar. + src_h, src_w = tensor.shape[-2:] + interp = ( + cv2.INTER_AREA + if self.model_w < src_w or self.model_h < src_h + else cv2.INTER_LINEAR + ) + tensor = np.stack( + [ + cv2.resize( + tensor[c], + (self.model_w, self.model_h), + interpolation=interp, + ) + for c in range(tensor.shape[0]) + ], + axis=0, + ).astype(np.float32) + + fusion_result = getattr(core, "last_fusion_result", None) + if isinstance(fusion_result, dict): + fusion_result = json.loads(json.dumps(fusion_result, default=str)) + # A prediction vive na resolução efetivamente entregue ao modelo. + fusion_result["target_size"] = [int(self.model_w), int(self.model_h)] + else: + fusion_result = None + + return tensor, fusion_result, mp_path + + def _infer_onnx(self, tensor: np.ndarray) -> Tuple[np.ndarray, float]: + x = tensor.astype(np.float32, copy=False) + if not self.normalization_embedded: + if self.mean is None or self.std is None: + raise RuntimeError("Normalização externa necessária, mas mean/std não estão carregados.") + x = (x - self.mean) / np.maximum(self.std, 1e-6) + + x = np.expand_dims(np.ascontiguousarray(x), axis=0).astype(np.float32) + t0 = time.perf_counter() + raw = self.session.run([self.selected_output_name], {self.input_name: x})[0] + infer_ms = (time.perf_counter() - t0) * 1000.0 + + if self.output_kind == "mask": + arr = np.asarray(raw) + while arr.ndim > 2 and arr.shape[0] == 1: + arr = arr[0] + if arr.ndim == 3 and arr.shape[0] == 1: + arr = arr[0] + if arr.ndim != 2: + raise RuntimeError( + f"Saída mask inesperada {self.selected_output_name}: {np.asarray(raw).shape}" + ) + pred = arr.astype(np.uint8) + if pred.shape != (self.model_h, self.model_w): + pred = cv2.resize( + pred, + (self.model_w, self.model_h), + interpolation=cv2.INTER_NEAREST, + ) + return pred, infer_ms + + logits = np.asarray(raw) + if logits.ndim == 4 and logits.shape[0] == 1: + logits = logits[0] + if logits.ndim != 3: + raise RuntimeError( + f"Saída logits inesperada {self.selected_output_name}: {np.asarray(raw).shape}" + ) + logits = _resize_logits_chw(logits, (self.model_h, self.model_w)) + probs = _softmax_np(logits, axis=0) + if self.head == "target" and probs.shape[0] > 1: + pred = (probs[1] >= float(self.target_threshold)).astype(np.uint8) + else: + pred = np.argmax(probs, axis=0).astype(np.uint8) + return pred, infer_ms + + def _infer_pt(self, tensor: np.ndarray) -> Tuple[np.ndarray, float]: + preds, _, infer_ms = self._pt_tester.infer(tensor) + pred = preds.get(self.head) + if pred is None: + raise RuntimeError( + f"Checkpoint não retornou head={self.head}. disponíveis={list(preds)}" + ) + pred = np.asarray(pred, dtype=np.uint8) + if pred.shape != (self.model_h, self.model_w): + pred = cv2.resize( + pred, + (self.model_w, self.model_h), + interpolation=cv2.INTER_NEAREST, + ) + return pred, float(infer_ms) + + def predict( + self, + bundle: SampleBundle, + rgb_size: Tuple[int, int], + module_params_override: Optional[Path], + ) -> PredictionArtifact: + def stat_sig(p: Path) -> Tuple[str, int, int]: + st = p.stat() + return str(p.resolve()), int(st.st_size), int(st.st_mtime_ns) + + cache_key = ( + stat_sig(bundle.meta_path), + tuple(stat_sig(bundle.bin_by_camera[c]) for c in REQUIRED_CAMERAS), + self.model_sha256, + self.head, + tuple(map(int, rgb_size)), + str(Path(module_params_override).resolve()) if module_params_override else "", + ) + if self._last_cache_key == cache_key and self._last_artifact is not None: + return self._last_artifact + + t0 = time.perf_counter() + tensor, fusion_result, mp_path = self._build_tensor(bundle, module_params_override) + tensor_ms = (time.perf_counter() - t0) * 1000.0 + + if not self._input_diag_printed: + stats = [] + for i, name in enumerate(self.input_channel_names): + ch = tensor[i] + stats.append( + f"{name}:min={float(np.nanmin(ch)):.4f}," + f"mean={float(np.nanmean(ch)):.4f}," + f"max={float(np.nanmax(ch)):.4f}" + ) + print( + "[SORTER_MODEL][INPUT] " + f"shape={tuple(tensor.shape)} norm_embedded={self.normalization_embedded} | " + + " | ".join(stats) + ) + self._input_diag_printed = True + + if self.model_path.suffix.lower() == ".onnx": + pred_ids, infer_ms = self._infer_onnx(tensor) + else: + pred_ids, infer_ms = self._infer_pt(tensor) + + vals, counts = np.unique(pred_ids, return_counts=True) + total = max(1, int(pred_ids.size)) + pred_summary = ", ".join( + f"{int(v)}={100.0*int(c)/total:.2f}%" for v, c in zip(vals, counts) + ) + print( + f"[SORTER_MODEL][PRED] sample={bundle.sample_id} head={self.head} " + f"ids={{ {pred_summary} }} infer={infer_ms:.1f}ms" + ) + + image_rgb, geom = prediction_ids_to_rgb_domain( + pred_ids, + fusion_result=fusion_result, + rgb_size=rgb_size, + head=self.head, + ) + + info = { + "schema": "weed_pair_sorter_prediction_v2", + "model_path": str(self.model_path), + "model_sha256": self.model_sha256, + "backend": self.backend, + "provider": self.active_provider, + "head": self.head, + "selected_output": self.selected_output_name, + "output_kind": self.output_kind, + "input_channel_names": list(self.input_channel_names), + "model_input_size": [int(self.model_w), int(self.model_h)], + "module_params_path": str(mp_path), + "target_threshold": float(self.target_threshold) if self.head == "target" else None, + "target_threshold_source": self.target_threshold_source if self.head == "target" else None, + "tensor_ms": float(tensor_ms), + "inference_ms": float(infer_ms), + "geometry": geom, + "saved_mask_encoding": ( + "semantic_rgb_labelmap_with_white_ignore" + if self.head == "semantic" + else f"{self.head}_binary_visual_mask_with_white_ignore" + ), + } + artifact = PredictionArtifact(image_rgb=image_rgb, info=info) + self._last_cache_key = cache_key + self._last_artifact = artifact + return artifact + + def add_label_to_image(img: Image.Image, title: str, subtitle: str = "") -> Image.Image: """Adiciona cabeçalho FORA do raster, sem esconder nenhum pixel da imagem.""" src = img.convert("RGB") @@ -1047,6 +2012,59 @@ def resize_keep_height(img: Image.Image, target_h: int) -> Image.Image: return img.resize((new_w, target_h), Image.BILINEAR) +def fit_image_within(img: Image.Image, max_width: int, max_height: int) -> Image.Image: + """Encaixa a imagem inteira na área disponível, preservando aspect ratio e sem crop.""" + src = img.convert("RGB") + w, h = src.size + max_width = max(64, int(max_width)) + max_height = max(64, int(max_height)) + if w <= 0 or h <= 0: + return src + + scale = min(max_width / float(w), max_height / float(h)) + new_w = max(1, int(round(w * scale))) + new_h = max(1, int(round(h * scale))) + if (new_w, new_h) == (w, h): + return src + + try: + resample = Image.Resampling.LANCZOS if scale < 1.0 else Image.Resampling.BILINEAR + except AttributeError: + resample = Image.LANCZOS if scale < 1.0 else Image.BILINEAR + return src.resize((new_w, new_h), resample) + + +def compose_prediction_overlay( + preview_rgb: Image.Image, + prediction_rgb: Optional[Image.Image], + overlay_pct: float, +) -> Image.Image: + """ + Mistura preview e prediction no MESMO domínio RGB canônico. + + 0% -> preview puro + 100% -> prediction pura (inclusive branco/ignore fora do crop) + + Como prediction_ids_to_rgb_domain() já reprojeta crop->RGB canônico, não há + resize geométrico independente aqui: apenas garantimos dimensões idênticas. + """ + preview = preview_rgb.convert("RGB") + if prediction_rgb is None: + return preview + + pred = prediction_rgb.convert("RGB") + if pred.size != preview.size: + # Fail-safe visual. A prediction correta deve chegar no mesmo domínio/tamanho. + pred = pred.resize(preview.size, Image.NEAREST) + + alpha = max(0.0, min(100.0, float(overlay_pct))) / 100.0 + if alpha <= 0.0: + return preview + if alpha >= 1.0: + return pred + return Image.blend(preview, pred, alpha) + + def compose_preview_grid(panels: List[Tuple[str, Image.Image, str]], display_height: int, max_width: int = 1800) -> Image.Image: """ Monta uma grade responsiva em PIL. @@ -1106,7 +2124,10 @@ def build_display_image( display_height: int, cache_root: Path, module_params_override: Optional[Path] = None, -) -> Tuple[ImageTk.PhotoImage, str, Path, dict, str]: + prediction_engine: Optional[PredictionEngine] = None, + display_width: int = 1800, + overlay_pct: float = 45.0, +) -> Tuple[ImageTk.PhotoImage, str, Path, dict, str, Optional[PredictionArtifact]]: panels: List[Tuple[str, Image.Image, str]] = [] primary_path, preview_report, preview_origin = ensure_canonical_preview( @@ -1127,6 +2148,37 @@ def build_display_image( panels.append(("Preview RGB para anotação", saved, origin_text)) status_extra = f"preview={preview_origin}" + prediction_artifact: Optional[PredictionArtifact] = None + if prediction_engine is not None: + try: + prediction_artifact = prediction_engine.predict( + bundle, + rgb_size=saved.size, + module_params_override=module_params_override, + ) + pinfo = prediction_artifact.info + subtitle = ( + f"{Path(str(pinfo.get('model_path', ''))).name} | " + f"head={pinfo.get('head')} | infer={float(pinfo.get('inference_ms', 0.0)):.1f}ms" + ) + if pinfo.get("target_threshold") is not None: + subtitle += f" | thr={float(pinfo['target_threshold']):.2f}" + # No modo full a prediction não vira um segundo painel: ela é + # sobreposta ao RGB pela trackbar. No modo de previews .bin mantemos + # o painel separado para diagnóstico multiespectral. + if show_bin_previews: + panels.append(( + f"Predição do modelo - {pinfo.get('head', '?')}", + Image.fromarray(prediction_artifact.image_rgb, mode="RGB"), + subtitle, + )) + status_extra += ( + f" | modelo={Path(str(pinfo.get('model_path', ''))).name}" + f" head={pinfo.get('head')}" + ) + except Exception as exc: + status_extra += f" | PREDIÇÃO INDISPONÍVEL: {exc}" + if show_bin_previews: errors = [] for cam_id in REQUIRED_CAMERAS: @@ -1141,8 +2193,34 @@ def build_display_image( if len(errors) > 2: status_extra += f" ; +{len(errors)-2}" - composed = compose_preview_grid(panels, display_height=display_height) - return ImageTk.PhotoImage(composed), status_extra, primary_path, preview_report, preview_origin + if not show_bin_previews: + pred_pil = ( + Image.fromarray(prediction_artifact.image_rgb, mode="RGB") + if prediction_artifact is not None + else None + ) + composed = compose_prediction_overlay(saved, pred_pil, overlay_pct) + composed = fit_image_within( + composed, + max_width=max(320, int(display_width)), + max_height=max(240, int(display_height)), + ) + if prediction_artifact is not None: + status_extra += f" | overlay={float(overlay_pct):.0f}%" + else: + composed = compose_preview_grid( + panels, + display_height=display_height, + max_width=max(640, int(display_width)), + ) + return ( + ImageTk.PhotoImage(composed), + status_extra, + primary_path, + preview_report, + preview_origin, + prediction_artifact, + ) # ============================================================ @@ -1236,6 +2314,7 @@ class SampleBundleSorterApp: display_height: int = 720, show_bin_previews: bool = False, module_params_override: Optional[Path] = None, + prediction_engine: Optional[PredictionEngine] = None, ): self.all_bundles = bundles self.labels = labels @@ -1244,10 +2323,16 @@ class SampleBundleSorterApp: self.display_height = display_height self.show_bin_previews_var_value = show_bin_previews self.module_params_override = Path(module_params_override) if module_params_override else None + self.prediction_engine = prediction_engine self.preview_cache_root = out_root / ".weeds_pair_sorter_cache" / "previews" self.current_primary_preview_path: Optional[Path] = None self.current_preview_report: Optional[dict] = None self.current_preview_origin: Optional[str] = None + self.current_prediction_artifact: Optional[PredictionArtifact] = None + self.current_prediction_key: Optional[str] = None + self.current_preview_pil: Optional[Image.Image] = None + self._resize_after_id = None + self._last_img_area = (0, 0) self.logger = ActionLogger(out_root) if resume: @@ -1262,10 +2347,18 @@ class SampleBundleSorterApp: (label_root / "previews").mkdir(parents=True, exist_ok=True) (label_root / "metas").mkdir(parents=True, exist_ok=True) (label_root / "masks").mkdir(parents=True, exist_ok=True) + if self.prediction_engine is not None: + (label_root / "predictions").mkdir(parents=True, exist_ok=True) self.root = tk.Tk() self.root.title("Weed Pair Sorter - RGB + RE + NIR") - self.root.geometry("1350x900") + self.root.geometry("1500x980") + # Em campo, usar o máximo de área útil sem entrar em fullscreen exclusivo. + # No Windows, zoomed mantém barra de título/atalhos do SO e maximiza a janela. + try: + self.root.state("zoomed") + except Exception: + pass self.root.bind("", self.on_key) self.top_frame = tk.Frame(self.root) @@ -1274,12 +2367,24 @@ class SampleBundleSorterApp: self.info_label = tk.Label(self.top_frame, text="", font=("Segoe UI", 11), anchor="w", justify="left") self.info_label.pack(side=tk.LEFT, padx=10, pady=6, fill=tk.X, expand=True) - self.legend_label = tk.Label(self.top_frame, text=self.build_legend_text(), font=("Segoe UI", 10), anchor="e") - self.legend_label.pack(side=tk.RIGHT, padx=10, pady=6) - self.opts_frame = tk.Frame(self.root) self.opts_frame.pack(side=tk.TOP, fill=tk.X) + # Legenda em linha própria: antes ela dividia a mesma faixa horizontal com + # info_label e sumia facilmente quando o texto da amostra ficava grande. + self.legend_label = tk.Label( + self.root, + text=self.build_legend_text(), + font=("Segoe UI", 10, "bold"), + anchor="w", + justify="left", + wraplength=1450, + bg="#f2f2f2", + padx=8, + pady=5, + ) + self.legend_label.pack(side=tk.TOP, fill=tk.X, padx=6, pady=(2, 4)) + self.show_bin_previews_var = tk.BooleanVar(value=self.show_bin_previews_var_value) tk.Checkbutton( self.opts_frame, @@ -1289,11 +2394,33 @@ class SampleBundleSorterApp: font=("Segoe UI", 10), ).pack(side=tk.LEFT, padx=10, pady=2) - self.img_frame = tk.Frame(self.root) + self.overlay_pct_var = tk.IntVar(value=45) + self.overlay_label = tk.Label( + self.opts_frame, + text="Overlay prediction (modo full):", + font=("Segoe UI", 10), + ) + self.overlay_label.pack(side=tk.LEFT, padx=(18, 4), pady=2) + self.overlay_scale = tk.Scale( + self.opts_frame, + from_=0, + to=100, + orient=tk.HORIZONTAL, + length=300, + resolution=1, + variable=self.overlay_pct_var, + command=self.on_overlay_change, + showvalue=True, + font=("Segoe UI", 9), + ) + self.overlay_scale.pack(side=tk.LEFT, padx=(0, 10), pady=0) + + self.img_frame = tk.Frame(self.root, bg="#222") self.img_frame.pack(side=tk.TOP, fill=tk.BOTH, expand=True) self.preview_label = tk.Label(self.img_frame, bg="#222") self.preview_label.pack(side=tk.TOP, expand=True, padx=6, pady=6) + self.img_frame.bind("", self.on_image_area_configure) self.status_var = tk.StringVar(value="Pronto.") self.status_label = tk.Label(self.root, textvariable=self.status_var, font=("Segoe UI", 10), anchor="w") @@ -1301,7 +2428,7 @@ class SampleBundleSorterApp: self.footer = tk.Label( self.root, - text="1..9/0 = labels | espaço/n/→ = pular | p/← = anterior | b = undo | v = liga/desliga previews .bin | q/Esc = sair", + text="1..9/0 = labels | espaço/n/→ = pular | p/← = anterior | b = undo | v = liga/desliga previews .bin | overlay 0%=RGB / 100%=prediction | q/Esc = sair", font=("Segoe UI", 10), ) self.footer.pack(side=tk.BOTTOM, fill=tk.X, pady=2) @@ -1315,6 +2442,79 @@ class SampleBundleSorterApp: parts.append(f"[{key}] {label}") return " | ".join(parts) + def _display_area_size(self) -> Tuple[int, int]: + """Área realmente disponível entre cabeçalho/legenda e rodapé.""" + try: + self.root.update_idletasks() + w = int(self.img_frame.winfo_width()) - 12 + h = int(self.img_frame.winfo_height()) - 12 + except Exception: + w, h = 0, 0 + + if w < 200: + try: + w = int(self.root.winfo_width()) - 20 + except Exception: + w = 1450 + if h < 160: + h = int(self.display_height) + return max(320, w), max(240, h) + + def _refresh_full_overlay_only(self): + """Atualiza somente a composição visual. Nunca roda o modelo novamente.""" + if self.show_bin_previews_var.get(): + return + if self.current_preview_pil is None: + return + + pred_pil = None + if self.current_prediction_artifact is not None: + pred_pil = Image.fromarray( + self.current_prediction_artifact.image_rgb, + mode="RGB", + ) + + composed = compose_prediction_overlay( + self.current_preview_pil, + pred_pil, + float(self.overlay_pct_var.get()), + ) + area_w, area_h = self._display_area_size() + composed = fit_image_within(composed, area_w, area_h) + self.preview_tk = ImageTk.PhotoImage(composed) + self.preview_label.configure(image=self.preview_tk) + + def on_overlay_change(self, _value=None): + # Tk.Scale chama isto muitas vezes durante o arraste. É propositalmente + # render-only para não disparar RawProcessor/TensorRT novamente. + self._refresh_full_overlay_only() + if self.prediction_engine is not None and not self.show_bin_previews_var.get(): + self.status_var.set( + f"Overlay prediction: {int(self.overlay_pct_var.get())}% " + "(0=RGB puro, 100=prediction pura)" + ) + + def on_image_area_configure(self, event): + new_area = (int(getattr(event, "width", 0)), int(getattr(event, "height", 0))) + if abs(new_area[0] - self._last_img_area[0]) < 8 and abs(new_area[1] - self._last_img_area[1]) < 8: + return + self._last_img_area = new_area + if self._resize_after_id is not None: + try: + self.root.after_cancel(self._resize_after_id) + except Exception: + pass + self._resize_after_id = self.root.after(120, self._on_resize_debounced) + + def _on_resize_debounced(self): + self._resize_after_id = None + if self.show_bin_previews_var.get(): + # A grade usa dimensões do frame e precisa ser recomposta. O PredictionEngine + # possui cache da amostra, então não refaz a inferência. + self.render() + else: + self._refresh_full_overlay_only() + def render(self): if not self.all_bundles: messagebox.showinfo("Fim", "Não há amostras para exibir.") @@ -1325,16 +2525,30 @@ class SampleBundleSorterApp: bundle = self.all_bundles[self.idx] try: - tk_img, status_extra, primary_path, preview_report, preview_origin = build_display_image( + area_w, area_h = self._display_area_size() + ( + tk_img, + status_extra, + primary_path, + preview_report, + preview_origin, + prediction_artifact, + ) = build_display_image( bundle, show_bin_previews=self.show_bin_previews_var.get() if hasattr(self, "show_bin_previews_var") else self.show_bin_previews_var_value, - display_height=self.display_height, + display_height=area_h, cache_root=self.preview_cache_root, module_params_override=self.module_params_override, + prediction_engine=self.prediction_engine, + display_width=area_w, + overlay_pct=float(self.overlay_pct_var.get()) if hasattr(self, "overlay_pct_var") else 45.0, ) self.current_primary_preview_path = primary_path self.current_preview_report = preview_report self.current_preview_origin = preview_origin + self.current_prediction_artifact = prediction_artifact + self.current_prediction_key = str(bundle.meta_path.resolve()) if prediction_artifact is not None else None + self.current_preview_pil = Image.open(primary_path).convert("RGB") self.preview_tk = tk_img self.preview_label.configure(image=self.preview_tk) except Exception as e: @@ -1374,11 +2588,46 @@ class SampleBundleSorterApp: out_base = choose_unique_output_base(label_root, bundle.sample_id) dst_preview = label_root / "previews" / f"{out_base}.png" dst_meta = label_root / "metas" / f"{out_base}.json" + dst_prediction = ( + label_root / "predictions" / f"{out_base}.png" + if self.prediction_engine is not None + else None + ) dst_by_camera = {cam: label_root / "bins" / f"{out_base}_{cam}.bin" for cam in REQUIRED_CAMERAS} dst_bins = [dst_by_camera[cam] for cam in REQUIRED_CAMERAS] src_bins = [bundle.bin_by_camera[cam] for cam in REQUIRED_CAMERAS] payload_files = {cam: dst_by_camera[cam].name for cam in REQUIRED_CAMERAS} + prediction_artifact = None + prediction_info = None + if self.prediction_engine is not None: + current_key = str(bundle.meta_path.resolve()) + if ( + self.current_prediction_artifact is not None + and self.current_prediction_key == current_key + ): + prediction_artifact = self.current_prediction_artifact + else: + with Image.open(primary_path) as _img: + rgb_size = _img.size + prediction_artifact = self.prediction_engine.predict( + bundle, + rgb_size=rgb_size, + module_params_override=self.module_params_override, + ) + + if dst_prediction is None: + raise RuntimeError("Destino prediction não foi criado.") + dst_prediction.parent.mkdir(parents=True, exist_ok=True) + Image.fromarray( + prediction_artifact.image_rgb, + mode="RGB", + ).save(dst_prediction, format="PNG", compress_level=1) + + prediction_info = dict(prediction_artifact.info) + prediction_info["saved_prediction_path"] = dst_prediction.name + prediction_info["saved_prediction_folder"] = "predictions" + # Preview é artefato derivado. Sempre copiamos a versão canônica gerada, # mesmo quando os RAW/meta são movidos. shutil.copy2(primary_path, dst_preview) @@ -1388,9 +2637,25 @@ class SampleBundleSorterApp: for src_bin, dst_bin in zip(src_bins, dst_bins): shutil.move(str(src_bin), str(dst_bin)) # Reescreve o meta já no destino com o contrato do preview. - write_destination_meta(dst_meta, dst_meta, dst_preview.name, preview_report, preview_origin, payload_files) + write_destination_meta( + dst_meta, + dst_meta, + dst_preview.name, + preview_report, + preview_origin, + payload_files, + prediction_info=prediction_info, + ) else: - write_destination_meta(bundle.meta_path, dst_meta, dst_preview.name, preview_report, preview_origin, payload_files) + write_destination_meta( + bundle.meta_path, + dst_meta, + dst_preview.name, + preview_report, + preview_origin, + payload_files, + prediction_info=prediction_info, + ) for src_bin, dst_bin in zip(src_bins, dst_bins): shutil.copy2(src_bin, dst_bin) @@ -1412,13 +2677,18 @@ class SampleBundleSorterApp: "meta_src": bundle.meta_path, "bin_srcs": list(src_bins), "preview_dst": dst_preview, + "prediction_dst": dst_prediction, "meta_dst": dst_meta, "bin_dsts": dst_bins, "moved": self.move, "index": self.idx, }) - self.status_var.set(f"{'Movido' if self.move else 'Copiado'} → '{label}': {bundle.sample_id} | preview canônico") + pred_txt = " | prediction salva" if prediction_info is not None else "" + self.status_var.set( + f"{'Movido' if self.move else 'Copiado'} → '{label}': " + f"{bundle.sample_id} | preview canônico{pred_txt}" + ) self.idx += 1 if self.idx >= len(self.all_bundles): messagebox.showinfo("Concluído", "Você chegou ao final da fila!") @@ -1441,6 +2711,7 @@ class SampleBundleSorterApp: meta_src = Path(last["meta_src"]) bin_srcs = [Path(p) for p in last["bin_srcs"]] preview_dst = Path(last["preview_dst"]) + prediction_dst = Path(last["prediction_dst"]) if last.get("prediction_dst") else None meta_dst = Path(last["meta_dst"]) bin_dsts = [Path(p) for p in last["bin_dsts"]] @@ -1463,6 +2734,8 @@ class SampleBundleSorterApp: if preview_dst.exists(): preview_dst.unlink() + if prediction_dst is not None and prediction_dst.exists(): + prediction_dst.unlink() self.logger.log( action="undo", @@ -1540,7 +2813,7 @@ class SetupWindow: def __init__(self): self.root = tk.Tk() self.root.title("Configurar - Weed Pair Sorter") - self.root.geometry("820x670") + self.root.geometry("940x900") frm_in = tk.LabelFrame(self.root, text="Pastas de entrada (raízes com subpastas de amostras)") frm_in.pack(fill=tk.BOTH, expand=False, padx=10, pady=8) @@ -1566,6 +2839,52 @@ class SetupWindow: tk.Entry(frm_mp, textvariable=self.module_params_var).pack(side=tk.LEFT, fill=tk.X, expand=True, padx=6, pady=6) tk.Button(frm_mp, text="Escolher JSON...", command=self.choose_module_params).pack(side=tk.RIGHT, padx=6, pady=6) + frm_model = tk.LabelFrame( + self.root, + text="Predição do modelo (opcional: .onnx ou .pt/.pth)", + ) + frm_model.pack(fill=tk.X, expand=False, padx=10, pady=8) + + row_model = tk.Frame(frm_model) + row_model.pack(fill=tk.X, padx=6, pady=3) + tk.Label(row_model, text="Modelo:", width=14, anchor="w").pack(side=tk.LEFT) + self.prediction_model_var = tk.StringVar(value="") + tk.Entry(row_model, textvariable=self.prediction_model_var).pack(side=tk.LEFT, fill=tk.X, expand=True) + tk.Button(row_model, text="Escolher...", command=self.choose_prediction_model).pack(side=tk.RIGHT, padx=(6, 0)) + + row_cfg = tk.Frame(frm_model) + row_cfg.pack(fill=tk.X, padx=6, pady=3) + tk.Label(row_cfg, text="Config:", width=14, anchor="w").pack(side=tk.LEFT) + self.prediction_config_var = tk.StringVar(value="config.json") + tk.Entry(row_cfg, textvariable=self.prediction_config_var).pack(side=tk.LEFT, fill=tk.X, expand=True) + tk.Button(row_cfg, text="Escolher...", command=self.choose_prediction_config).pack(side=tk.RIGHT, padx=(6, 0)) + + row_test = tk.Frame(frm_model) + row_test.pack(fill=tk.X, padx=6, pady=3) + tk.Label(row_test, text="Test script (.pt):", width=14, anchor="w").pack(side=tk.LEFT) + self.prediction_test_script_var = tk.StringVar(value="_9_test_multihead_v2.py") + tk.Entry(row_test, textvariable=self.prediction_test_script_var).pack(side=tk.LEFT, fill=tk.X, expand=True) + tk.Button(row_test, text="Escolher...", command=self.choose_prediction_test_script).pack(side=tk.RIGHT, padx=(6, 0)) + + row_adv = tk.Frame(frm_model) + row_adv.pack(fill=tk.X, padx=6, pady=3) + tk.Label(row_adv, text="Head:", width=7, anchor="w").pack(side=tk.LEFT) + self.prediction_head_var = tk.StringVar(value="semantic") + tk.OptionMenu(row_adv, self.prediction_head_var, "semantic", "target", "vegetation", "cana").pack(side=tk.LEFT) + tk.Label(row_adv, text="Provider:", padx=8).pack(side=tk.LEFT) + self.prediction_provider_var = tk.StringVar(value="tensorrt") + tk.OptionMenu(row_adv, self.prediction_provider_var, "tensorrt", "cuda", "cpu").pack(side=tk.LEFT) + tk.Label(row_adv, text="Target threshold:", padx=8).pack(side=tk.LEFT) + self.prediction_threshold_var = tk.StringVar(value="auto") + tk.Entry(row_adv, textvariable=self.prediction_threshold_var, width=7).pack(side=tk.LEFT) + + row_norm = tk.Frame(frm_model) + row_norm.pack(fill=tk.X, padx=6, pady=3) + tk.Label(row_norm, text="Norm stats:", width=14, anchor="w").pack(side=tk.LEFT) + self.prediction_norm_stats_var = tk.StringVar(value="") + tk.Entry(row_norm, textvariable=self.prediction_norm_stats_var).pack(side=tk.LEFT, fill=tk.X, expand=True) + tk.Button(row_norm, text="Escolher...", command=self.choose_prediction_norm_stats).pack(side=tk.RIGHT, padx=(6, 0)) + frm_labels = tk.LabelFrame(self.root, text="Labels (classes) separadas por vírgula") frm_labels.pack(fill=tk.X, expand=False, padx=10, pady=8) self.labels_var = tk.StringVar(value="chao,chao_cana,chao_erva,chao_cana_erva,cana,cana_erva,erva") @@ -1627,6 +2946,43 @@ class SetupWindow: if p: self.module_params_var.set(p) + def choose_prediction_model(self): + p = filedialog.askopenfilename( + title="Selecione ONNX ou checkpoint PyTorch", + filetypes=[ + ("Modelos", "*.onnx *.pt *.pth"), + ("ONNX", "*.onnx"), + ("PyTorch", "*.pt *.pth"), + ("Todos", "*.*"), + ], + ) + if p: + self.prediction_model_var.set(p) + + def choose_prediction_config(self): + p = filedialog.askopenfilename( + title="Selecione config.json do treinamento/modelo", + filetypes=[("JSON", "*.json"), ("Todos", "*.*")], + ) + if p: + self.prediction_config_var.set(p) + + def choose_prediction_test_script(self): + p = filedialog.askopenfilename( + title="Selecione _9_test_multihead_v2.py", + filetypes=[("Python", "*.py"), ("Todos", "*.*")], + ) + if p: + self.prediction_test_script_var.set(p) + + def choose_prediction_norm_stats(self): + p = filedialog.askopenfilename( + title="Selecione norm_stats.json (opcional)", + filetypes=[("JSON", "*.json"), ("Todos", "*.*")], + ) + if p: + self.prediction_norm_stats_var.set(p) + def start(self): inputs = [self.inputs_listbox.get(i) for i in range(self.inputs_listbox.size())] out_root = self.out_root_var.get().strip() @@ -1652,6 +3008,13 @@ class SetupWindow: self.height_var.get(), self.show_bin_previews_var.get(), Path(self.module_params_var.get().strip()) if self.module_params_var.get().strip() else None, + Path(self.prediction_model_var.get().strip()) if self.prediction_model_var.get().strip() else None, + Path(self.prediction_config_var.get().strip()) if self.prediction_config_var.get().strip() else None, + Path(self.prediction_test_script_var.get().strip()) if self.prediction_test_script_var.get().strip() else None, + Path(self.prediction_norm_stats_var.get().strip()) if self.prediction_norm_stats_var.get().strip() else None, + self.prediction_provider_var.get().strip(), + self.prediction_head_var.get().strip(), + self.prediction_threshold_var.get().strip() or "auto", ) self.root.destroy() @@ -1666,7 +3029,42 @@ def run_with_gui_setup(): if not res: return - input_folders, labels, out_root, move, resume, height, show_bin_previews, module_params_override = res + ( + input_folders, + labels, + out_root, + move, + resume, + height, + show_bin_previews, + module_params_override, + prediction_model, + prediction_config, + prediction_test_script, + prediction_norm_stats, + prediction_provider, + prediction_head, + prediction_threshold, + ) = res + + prediction_engine = None + if prediction_model is not None: + try: + prediction_engine = PredictionEngine( + model_path=prediction_model, + config_path=prediction_config, + test_script_path=prediction_test_script, + norm_stats_path=prediction_norm_stats, + provider=prediction_provider, + head=prediction_head, + target_threshold=prediction_threshold, + ) + except Exception as exc: + messagebox.showerror( + "Erro ao carregar modelo", + f"Não foi possível iniciar a predição opcional:\n\n{exc}", + ) + return bundles = collect_all_bundles(input_folders) if not bundles: @@ -1687,6 +3085,7 @@ def run_with_gui_setup(): display_height=height, show_bin_previews=show_bin_previews, module_params_override=module_params_override, + prediction_engine=prediction_engine, ) app.run() @@ -1703,6 +3102,13 @@ def main(): parser.add_argument("--display-height", type=int, default=760, help="Altura total de exibição do preview/grid (px)") parser.add_argument("--show-bin-previews", action="store_true", help="Inicia mostrando CAM_A/CAM_B/CAM_C reconstruídas dos .bin") parser.add_argument("--module-params", default=None, help="Fallback/override do module_params/modelmp para reconstruir preview canônico") + parser.add_argument("--prediction-model", default=None, help="Modelo opcional .onnx/.pt/.pth para mostrar/salvar prediction") + parser.add_argument("--prediction-config", default="config.json", help="config.json do modelo; obrigatório para .pt e fallback do ONNX") + parser.add_argument("--prediction-test-script", default="_9_test_multihead_v2.py", help="Script de teste usado para reconstruir checkpoint .pt") + parser.add_argument("--prediction-norm-stats", default=None, help="norm_stats.json opcional; auto quando omitido") + parser.add_argument("--prediction-provider", default="tensorrt", choices=["tensorrt", "cuda", "cpu"]) + parser.add_argument("--prediction-head", default="semantic", choices=["semantic", "target", "vegetation", "cana"]) + parser.add_argument("--prediction-target-threshold", default="auto", help="auto ou valor 0..1; usado quando head=target e o modelo retorna logits") parser.add_argument("--no-gui-setup", action="store_true", help="Não abrir a GUI de setup") args = parser.parse_args() @@ -1719,6 +3125,18 @@ def main(): labels = args.labels out_root = Path(args.out_root) + prediction_engine = None + if args.prediction_model: + prediction_engine = PredictionEngine( + model_path=Path(args.prediction_model), + config_path=Path(args.prediction_config) if args.prediction_config else None, + test_script_path=Path(args.prediction_test_script) if args.prediction_test_script else None, + norm_stats_path=Path(args.prediction_norm_stats) if args.prediction_norm_stats else None, + provider=args.prediction_provider, + head=args.prediction_head, + target_threshold=args.prediction_target_threshold, + ) + bundles = collect_all_bundles(input_folders) if not bundles: print( @@ -1737,6 +3155,7 @@ def main(): display_height=args.display_height, show_bin_previews=args.show_bin_previews, module_params_override=Path(args.module_params) if args.module_params else None, + prediction_engine=prediction_engine, ) app.run() diff --git a/Python/OAK/datasets/oak-fcc-3/_9_test_multihead_v2.py b/Python/OAK/datasets/oak-fcc-3/_9_test_multihead_v2.py index 2f1a526d5..c4f0628fd 100644 --- a/Python/OAK/datasets/oak-fcc-3/_9_test_multihead_v2.py +++ b/Python/OAK/datasets/oak-fcc-3/_9_test_multihead_v2.py @@ -27,6 +27,15 @@ Exemplo: python .\\_9_test_multihead.py --config config.json --split_folder val --ckpt backup\segformer_b1\test_multi\stacked_raw5_multihead\best_score.pt + +python .\_9_test_multihead_v2.py ` + --config .\config_bench_ar0234.json ` + --test_folder "C:\Users\USER\Desktop\fotos_multiespectrais\empyreo\chao_cana_erva" ` + --ckpt .\backup\segformer_b1\2026_09_18\stacked_raw5\best_score.pt ` + --norm_stats .\dataset\960x600\group\norm_stats.json ` + --runtime_mode all ` + --target_threshold auto + Controles: D / seta direita : próxima amostra A / seta esquerda: amostra anterior diff --git a/Python/OAK/datasets/oak-fcc-3/audit/review_dataset_multihead_v2.py b/Python/OAK/datasets/oak-fcc-3/audit/review_dataset_multihead_v2.py index 169e1e7be..784a84381 100644 --- a/Python/OAK/datasets/oak-fcc-3/audit/review_dataset_multihead_v2.py +++ b/Python/OAK/datasets/oak-fcc-3/audit/review_dataset_multihead_v2.py @@ -1629,6 +1629,10 @@ def main() -> None: raw_root: Optional[Path] = None raw_source_map: Dict[Tuple[str, str], dict] = {} raw_source_resolution: Dict[Tuple[str, str], str] = {} + # Amostras normalizadas continuam válidas para auditoria mesmo quando o + # bundle RAW original está incompleto/ausente. Nesses casos, só bloqueamos + # a exportação física RAW-SAFE daquela amostra e registramos o motivo. + raw_source_failures: Dict[Tuple[str, str], str] = {} if not args.report_only: raw_root = ( @@ -1645,13 +1649,22 @@ def main() -> None: raw_by_key, raw_by_base = index_raw_review_sources(raw_root) fallback_count = 0 for sample in samples: - rec, mode = resolve_raw_review_source( - sample, - raw_by_key, - raw_by_base, - dataset_path, - ) key = (str(sample.group), str(sample.base)) + try: + rec, mode = resolve_raw_review_source( + sample, + raw_by_key, + raw_by_base, + dataset_path, + ) + except RuntimeError as exc: + raw_source_resolution[key] = "unavailable" + raw_source_failures[key] = str(exc) + print( + f"[RAW-SAFE][SKIP-PHYSICAL] {sample.group}/{sample.base} | {exc}" + ) + continue + raw_source_map[key] = rec raw_source_resolution[key] = mode if mode != "same_group": @@ -1659,6 +1672,7 @@ def main() -> None: print( f"[RAW-SAFE] raw_root={raw_root} | resolvidos={len(raw_source_map)}/{len(samples)} " + f"| indisponiveis={len(raw_source_failures)} " f"| fallback_base_unico={fallback_count}" ) @@ -2029,6 +2043,8 @@ def main() -> None: domain_manifest_rows: List[dict] = [] domain_manifest_path = review_group_root / "review_domain_manifest.csv" domain_readme_path = review_group_root / "README_DOMAIN_CONTRACT.txt" + physical_exported_count = 0 + physical_export_skipped_count = 0 if not args.report_only: print(f"\n[COPY RAW-SAFE] candidatos finais: {len(selected_rows)}") @@ -2038,6 +2054,23 @@ def main() -> None: key = (str(row["group"]), str(row["base"])) sample = sample_by_key[key] + # A auditoria/relatório desta amostra continua válido. Só não há + # como gerar raw_previews/raw_masks/predictions_raw com segurança. + raw_source = raw_source_map.get(key) + if raw_source is None: + physical_export_skipped_count += 1 + row["physical_export_skipped"] = 1 + row["physical_export_skip_reason"] = raw_source_failures.get( + key, + "Fonte RAW indisponível para exportação física RAW-SAFE", + ) + print( + f"[COPY RAW-SAFE][SKIP] {sample.group}/{sample.base} | " + f"{row['physical_export_skip_reason']}" + ) + cache_for_copy.pop((sample.group, sample.base), None) + continue + cached = cache_for_copy.get((sample.group, sample.base)) if cached is not None and cached[0].group == sample.group: _, chw, gt_sem, pred_sem = cached @@ -2057,26 +2090,42 @@ def main() -> None: pred_sem = preds["semantic"] gt_sem = resize_ids(base.load_mask(sample.masks["semantic"]), pred_sem.shape[:2]) - raw_source = raw_source_map.get((str(sample.group), str(sample.base))) - if raw_source is None: - raise RuntimeError( - f"Fonte RAW desapareceu durante exportação: {sample.group}/{sample.base}" + try: + domain_row = save_review_item( + row=row, + sample=sample, + chw=chw, + gt_sem=gt_sem, + pred_sem=pred_sem, + raw_source=raw_source, + review_group_root=review_group_root, + base_module=base, + semantic_cmap=semantic_cmap, + ignore_id=ignore_id, + input_channel_names=input_channel_names, + save_panels=args.save_panels, ) + except RuntimeError as exc: + # Um bundle RAW existente, mas geometricamente inválido/corrompido, + # também não deve abortar milhares de amostras já auditáveis. + physical_export_skipped_count += 1 + row["physical_export_skipped"] = 1 + row["physical_export_skip_reason"] = str(exc) + print( + f"[COPY RAW-SAFE][SKIP] {sample.group}/{sample.base} | {exc}" + ) + cache_for_copy.pop((sample.group, sample.base), None) + del chw, gt_sem, pred_sem + if 'preds' in locals(): + try: + del preds + except Exception: + pass + continue - domain_row = save_review_item( - row=row, - sample=sample, - chw=chw, - gt_sem=gt_sem, - pred_sem=pred_sem, - raw_source=raw_source, - review_group_root=review_group_root, - base_module=base, - semantic_cmap=semantic_cmap, - ignore_id=ignore_id, - input_channel_names=input_channel_names, - save_panels=args.save_panels, - ) + physical_exported_count += 1 + row["physical_export_skipped"] = 0 + row["physical_export_skip_reason"] = "" domain_row["raw_source_resolution"] = raw_source_resolution.get( (str(sample.group), str(sample.base)), "" ) @@ -2249,6 +2298,10 @@ def main() -> None: "samples": len(rows), "selected_for_review": len(selected_rows), "selected_pct": float(100.0 * len(selected_rows) / max(len(rows), 1)), + "raw_source_available": len(raw_source_map) if not args.report_only else None, + "raw_source_unavailable": len(raw_source_failures) if not args.report_only else None, + "physical_exported": physical_exported_count if not args.report_only else 0, + "physical_export_skipped": physical_export_skipped_count if not args.report_only else 0, "export": { "mode": str(args.export_mode), "min_suspicion_pct": float(args.min_suspicion_pct), @@ -2358,6 +2411,10 @@ def main() -> None: print(f"Summary : {summary_path}") if not args.report_only: print(f"Revisão física RAW-SAFE : {review_group_root}") + print(f"RAW disponíveis : {len(raw_source_map)}") + print(f"RAW indisponíveis : {len(raw_source_failures)}") + print(f"Exportados fisicamente : {physical_exported_count}") + print(f"Pulados na exportação RAW : {physical_export_skipped_count}") print(f"Domain manifest : {domain_manifest_path}") print(f"Domain contract : {domain_readme_path}") diff --git a/Python/OAK/datasets/oak-fcc-3/calibration/mp_ar0234/module_params.json b/Python/OAK/datasets/oak-fcc-3/calibration/mp_ar0234/module_params.json index 46c1dc6c7..acd41bf25 100644 --- a/Python/OAK/datasets/oak-fcc-3/calibration/mp_ar0234/module_params.json +++ b/Python/OAK/datasets/oak-fcc-3/calibration/mp_ar0234/module_params.json @@ -721,7 +721,7 @@ } }, "flatfield_config": { - "enabled": true, + "enabled": false, "npz_file": "flatfield_maps_v1.npz", "apply_before_fusion": true, "apply_after_decode": true, diff --git a/Python/OAK/datasets/oak-fcc-3/config_bench_ar0234.json b/Python/OAK/datasets/oak-fcc-3/config_bench_ar0234.json new file mode 100644 index 000000000..98f3eb1b8 --- /dev/null +++ b/Python/OAK/datasets/oak-fcc-3/config_bench_ar0234.json @@ -0,0 +1,216 @@ +{ + "camera": "oak-fcc-3", + "modelo": "segformer_b1", + "model_name": "2026_09_08", + "main_class_name": "erva", + "es_classes": "", + "model_to_use": "geral", + "raw_size": [1280, 800], + "resolucao": [960, 600], + "roi_inicio": 0.0, + "roi_tamanho": 1.0, + "shaves": 3, + "source_channels": ["R", "G", "B", "RE", "NIR"], + "channels": 5, + "input_channels": ["R", "G", "B", "RE", "NIR"], + "derived_channels": { + "epsilon": 1e-6, + "clip_min": -1.0, + "clip_max": 1.0 + }, + "backbone": "nvidia/mit-b1", + "fusion_mode": "stacked", + "stats_source_tag": "stacked_raw5", + "module_params_json": "calibration/mp_ar0234/module_params.json", + "ckpt_test": "best_operational", + "multi_head": true, + "heads": { + "semantic": { + "enabled": true, + "type": "multiclass", + "num_classes": 3, + "mask_dir": "masks", + "classes": {"chao": 0, "cana": 1, "erva": 2}, + "ignore_index": 255, + "loss_weight": 0.20 + }, + "vegetation": { + "enabled": true, + "type": "binary", + "num_classes": 2, + "mask_dir": "masks_vegetation", + "classes": {"background": 0, "vegetation": 1}, + "ignore_index": 255, + "loss_weight": 0.15 + }, + "cana": { + "enabled": true, + "type": "binary", + "num_classes": 2, + "mask_dir": "masks_cana", + "classes": {"not_cana": 0, "cana": 1}, + "ignore_index": 255, + "loss_weight": 0.35 + }, + "target": { + "enabled": true, + "type": "binary", + "num_classes": 2, + "mask_dir": "__derived_target__", + "classes": {"background": 0, "target": 1}, + "ignore_index": 255, + "loss_weight": 0.30, + "derived": true + } + }, + "training_v2": { + "model": { + "decoder_mode": "shared_light", + "spectral_input_init": "zero_extra" + }, + "augmentation": { + "enabled": true, + "horizontal_flip_p": 0.5, + "vertical_flip_p": 0.0, + "affine_p": 0.7, + "rotate_deg": 5.0, + "scale_min": 0.9, + "scale_max": 1.1, + "translate_frac": 0.04, + "crop_p": 0.45, + "crop_scale_min": 0.7, + "crop_scale_max": 1.0, + "crop_focus_target_p": 0.55, + "crop_focus_cana_p": 0.25, + "global_gain_p": 0.35, + "global_gain_min": 0.92, + "global_gain_max": 1.08, + "band_gain_p": 0.25, + "band_gain_min": 0.96, + "band_gain_max": 1.04, + "rgb_gamma_p": 0.2, + "rgb_gamma_min": 0.94, + "rgb_gamma_max": 1.06, + "noise_p": 0.2, + "noise_sigma_min": 0.001, + "noise_sigma_max": 0.008, + "blur_p": 0.12, + "blur_kernel": 3, + "sensor_channel_dropout_p": 0.0, + "clip_physical": true + }, + "sampler": { + "mode": "diverse", + "samples_per_epoch": 0, + "tiny_target_pct": 0.005, + "small_target_pct": 0.02, + "medium_target_pct": 0.1 + }, + "class_weighting": { + "method": "log_inverse", + "log_offset": 1.02, + "power": 0.5, + "min_weight": 0.25, + "max_weight": 4.0 + }, + "loss": { + "dice_reduction": "per_image", + "dice_smooth": 1.0, + "boundary_weight": 0.0, + "ohem_ratio": 0.0, + "safety": { + "enabled": true, + "weight": 0.08, + "cana_weight": 1.0, + "ground_weight": 0.2 + } + }, + "optimizer": { + "encoder_lr": null, + "patch_lr_mult": 2.0, + "heads_lr_mult": 5.0, + "weight_decay": null, + "no_decay_bias": true, + "no_decay_norm": true, + "betas": [ + 0.9, + 0.999 + ], + "eps": 1e-08 + }, + "scheduler": { + "mode": "poly", + "warmup_ratio": 0.05, + "warmup_start_factor": 0.1, + "poly_power": 1.0, + "min_lr_ratio": 0.02 + }, + "optimization": { + "grad_clip_norm": 1.0, + "matmul_precision": "high", + "cudnn_benchmark": true, + "persistent_workers": true, + "prefetch_factor": 2 + }, + "target_distillation": { + "enabled": false, + "mode": "cross_head", + "start_epoch": 8, + "rampup_epochs": 12, + "hard_weight": 0.75, + "distill_weight": 0.25, + "teacher_confidence_min": 0.6, + "detach_teacher": true, + "w_sem_erva": 0.45, + "w_veg_not_cana": 0.35, + "w_veg_suppressed": 0.2, + "cana_suppression_power": 1.5 + }, + "metrics": { + "target_thresholds": [ + 0.3, + 0.4, + 0.5, + 0.6, + 0.7, + 0.8, + 0.9 + ], + "ece_bins": 15, + "scenario_metrics": true, + "group_metrics": true, + "rich_train_metrics": false, + "operational_threshold": { + "max_cana_spray_rate": 0.02, + "max_ground_spray_rate": 0.03, + "max_weed_miss_rate": 0.2, + "score_weights": { + "target_iou": 0.35, + "target_f1": 0.2, + "cana_safety": 0.25, + "ground_safety": 0.1, + "weed_recall": 0.1 + } + } + }, + "selection_score": { + "target_iou": 0.4, + "cana_iou": 0.2, + "target_f1": 0.1, + "vegetation_miou": 0.1, + "semantic_miou": 0.05, + "cana_safety": 0.15 + }, + "checkpoint": { + "early_stop_min_delta": 0.0005, + "save_best_safety": true, + "save_best_legacy": true, + "save_best_operational": true + }, + "data": { + "validate_npy_content": true, + "skip_corrupt_samples": true, + "max_corrupt_fraction": 0.005 + } + } +} \ No newline at end of file diff --git a/Python/OAK/datasets/oak-fcc-3/core/raw_processor_core.py b/Python/OAK/datasets/oak-fcc-3/core/raw_processor_core.py index 3aa9385dd..3c7dd7144 100644 --- a/Python/OAK/datasets/oak-fcc-3/core/raw_processor_core.py +++ b/Python/OAK/datasets/oak-fcc-3/core/raw_processor_core.py @@ -872,7 +872,7 @@ class RawProcessorCore: # OAK_CORE_PERF_LOG_INTERVAL_S=1.0 # OAK_CORE_SHAPES_LOG_INTERVAL_S=5.0 self.core_perf_debug = str( - os.getenv("OAK_CORE_PERF_DEBUG", "1") + os.getenv("OAK_CORE_PERF_DEBUG", "0") ).strip().lower() not in ("0", "false", "no", "off") try: self.core_perf_log_interval_s = max(0.2, float(