# raw_segformer_service.py # Serviço para SegFormer B3 treinado em RAW com 4 ou 5 canais. # # Responsabilidades: # - Carregar labelmap e checkpoint (.pt) # - Montar o modelo SegFormer com N canais de entrada (4 ou 5) # - Normalizar entrada com normalize_raw (mesma do treino) # - Fazer inferência e devolver pred_ids [H,W] # - Gerar preview RGB simples a partir da entrada import os import json from contextlib import nullcontext from typing import List, Optional, Tuple import threading import time import numpy as np import torch import torch.nn as nn from torch import amp from torch.utils.data import Dataset import torch.nn.functional as F from transformers import SegformerForSemanticSegmentation from utils import (carregar_labelmap_completo, converter_mask_ids_para_bgr) def build_raw_segformer_model( num_classes: int, channels: int, backbone: str = "nvidia/segformer-b3-finetuned-ade-512-512", device: torch.device | None = None, ckpt_path: str | None = None, strict: bool = True, ) -> SegformerForSemanticSegmentation: if device is None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = SegformerForSemanticSegmentation.from_pretrained( backbone, num_labels=num_classes, ignore_mismatched_sizes=True, use_safetensors=True, ) # patch pra 4 ou 5 canais patch_segformer_input_channels(model, in_ch=channels) if ckpt_path is not None: ckpt = torch.load(ckpt_path, map_location="cpu") if "model" not in ckpt: raise RuntimeError(f"Checkpoint {ckpt_path} não contém chave 'model'.") model.load_state_dict(ckpt["model"], strict=strict) model.to(device) return model class RawSegformerService: """ Service para trabalhar com o modelo SegFormer RAW (4 ou 5 canais). Exemplo de uso: svc = RawSegformerService( config_path="config.json", ckpt_path="gal5000/backup/segformer_b3/dif/rawx/best_miou.pt", device=torch.device("cuda"), ) # raw_np: np.ndarray (C,H,W) float32 em 0..1 (C=4 ou 5) pred_ids, rgb, pred_rgb, overlay = svc.infer_and_preview(raw_np) """ def __init__( self, config_path: str, ckpt_path: str | None = None, labelmap_path: str | None = None, device: torch.device | None = None, use_amp: bool = True, mean=None, std=None, load_model=True ): # device if device is None: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.device = device self.use_amp = use_amp and (self.device.type == "cuda") # Setups "free speed" torch.backends.cudnn.benchmark = True if device.type == "cuda": torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True # mean/std FIXOS (recomendado) # mean/std devem ser listas/np arrays com tamanho C self._norm_mean = None self._norm_std = None if mean is not None and std is not None: self.set_norm_stats(mean, std) # Buffers self._buf_shape = None self._cpu_pinned = None # (1,C,H,W) float32 pinned self._gpu_input = None # (1,C,H,W) float32 no device # carrega config with open(config_path, "r") as f: cfg = json.load(f) self.cfg = cfg self.MODELO = cfg["camera"] # ex: "gal5000" self.MODEL_NAME = cfg["model_name"] # ex: "segformer_b3" self.modelo_folder = cfg["modelo"] # ex: "dif" self.res_w, self.res_h = cfg["resolucao"] # [W,H] self.backbone = cfg["backbone"] self._preview_lock = threading.Lock() self._preview_busy = False self._last_preview = None self._last_preview_ts = 0.0 self._last_preview_dt = 0.0 # Buffers reutilizáveis (alocados sob demanda) self._buf_hw = None # (H, W) atual self._raw_buf4 = None # (4, H, W) self._raw_buf5 = None # (5, H, W) # Buffers auxiliares p/ NDVI (H, W) self._ndvi_buf = None # ndvi em [-1,1] self._den_buf = None # denominador (ir + r + eps) # novos parâmetros self.channels = int(cfg.get("channels", 4)) # 4 ou 5 self.use_ndvi = bool(cfg.get("use_ndvi", False)) if self.channels not in (4, 5): raise ValueError(f"[svc] 'channels' inválido no config: {self.channels} (esperado 4 ou 5)") if self.channels == 5 and not self.use_ndvi: raise ValueError("[svc] channels=5 sem NDVI não faz sentido no design atual. " "Use channels=4 sem NDVI ou channels=5 com NDVI.") # labelmap if labelmap_path is None: dataset_path = os.path.join(self.MODELO, "dataset") labelmap_path = os.path.join(dataset_path, "labelmap.txt") self.labelmap_path = labelmap_path _, colormap_rgb, classes, ignore_rgb = carregar_labelmap_completo(labelmap_path) self.colormap_rgb = colormap_rgb self.classes = classes self.ignore_id = _infer_ignore_id(ignore_rgb, default_id=255) self.num_classes = len(classes) print(f"[svc] classes={self.classes}") print(f"[svc] ignore_id={self.ignore_id}") print(f"[svc] num_classes={self.num_classes}") print(f"[svc] channels={self.channels} use_ndvi={self.use_ndvi}") if load_model: # checkpoint if ckpt_path is None: # igual ao _9_test_segformer_b3_raw.py # save_path = MODELO/backup/modelo/model_name/rawx save_path = os.path.join(self.MODELO, "backup", self.modelo_folder, self.MODEL_NAME, f"raw{self.channels}") ckpt_path = os.path.join(save_path, "best_miou.pt") if not os.path.isfile(ckpt_path): raise FileNotFoundError(f"Checkpoint não encontrado: {ckpt_path}") self.ckpt_path = ckpt_path print(f"[svc] ckpt_path={ckpt_path}") # modelo self.model = self._load_model(backbone=self.backbone) self.model.eval() def _ensure_buffers(self, H: int, W: int): """Garante que os buffers internos têm shape compatível com (H,W).""" if self._buf_hw == (H, W): return self._buf_hw = (H, W) # Buffers principais self._raw_buf4 = np.empty((4, H, W), dtype=np.float32) self._raw_buf5 = np.empty((5, H, W), dtype=np.float32) # Buffers auxiliares para NDVI self._ndvi_buf = np.empty((H, W), dtype=np.float32) self._den_buf = np.empty((H, W), dtype=np.float32) def set_norm_stats(self, mean, std): mean = torch.tensor(mean, dtype=torch.float32, device=self.device).view(1, -1, 1, 1) std = torch.tensor(std, dtype=torch.float32, device=self.device).view(1, -1, 1, 1).clamp_min(1e-6) self._norm_mean = mean self._norm_std = std def prepare_infer_buffers(self, C, H, W): shape = (1, C, H, W) if self._buf_shape == shape: return self._buf_shape = shape # pinned CPU buffer (acelera H2D) self._cpu_pinned = torch.empty(shape, dtype=torch.float32, pin_memory=True) # GPU input buffer self._gpu_input = torch.empty(shape, dtype=torch.float32, device=self.device) # ------------------------------------------------- # Carregamento do modelo # ------------------------------------------------- def _load_model(self, backbone: str) -> torch.nn.Module: model = build_raw_segformer_model( num_classes=self.num_classes, channels=self.channels, backbone=backbone, device=self.device, ckpt_path=self.ckpt_path, strict=True, ) print(f"[svc] modelo carregado: {self.ckpt_path}") return model # ------------------------------------------------- # Pré-processamento / preview # ------------------------------------------------- def build_raw_input(self, r: np.ndarray, g: np.ndarray, ir: np.ndarray, b: np.ndarray) -> np.ndarray: """ Versão com buffers reutilizáveis. Entradas: r,g,ir,b: (H,W) float32 em 0..1 (idealmente já float32) Saída: (C,H,W) float32 conforme: - channels=4, use_ndvi=False -> [R,G,B,IR] - channels=4, use_ndvi=True -> [R,G,B,NDVI01] - channels=5, use_ndvi=True -> [R,G,B,IR,NDVI01] Observação: Retorna um buffer interno reutilizado. Se você guardar a referência e chamar de novo, o conteúdo será sobrescrito. (No loop de tempo real isso é perfeito.) """ if self.channels not in (4, 5): raise RuntimeError(f"channels inválido em build_raw_input: {self.channels}") # Garantir float32 sem copiar se já estiver r32 = np.asarray(r, dtype=np.float32) g32 = np.asarray(g, dtype=np.float32) ir32 = np.asarray(ir, dtype=np.float32) b32 = np.asarray(b, dtype=np.float32) H, W = r32.shape self._ensure_buffers(H, W) # Seleciona buffer principal raw = self._raw_buf5 if self.channels == 5 else self._raw_buf4 if not self.use_ndvi: if self.channels != 4: raise RuntimeError("channels=5 sem NDVI não está definido em build_raw_input.") # [R,G,B,IR] raw[0, :, :] = r32 raw[1, :, :] = g32 raw[2, :, :] = b32 raw[3, :, :] = ir32 return raw # use_ndvi=True eps = np.float32(1e-6) # Preenche base RGB (e IR se channels=5) raw[0, :, :] = r32 raw[1, :, :] = g32 raw[2, :, :] = b32 if self.channels == 5: raw[3, :, :] = ir32 # Calcula NDVI usando buffers auxiliares: # den = ir + r + eps np.add(ir32, r32, out=self._den_buf) # den = ir + r self._den_buf += eps # den += eps # ndvi = (ir - r) / den np.subtract(ir32, r32, out=self._ndvi_buf) # ndvi_buf = ir - r np.divide(self._ndvi_buf, self._den_buf, out=self._ndvi_buf) # ndvi_buf /= den # clip [-1,1] np.clip(self._ndvi_buf, -1.0, 1.0, out=self._ndvi_buf) # ndvi01 = (ndvi + 1) * 0.5 self._ndvi_buf += np.float32(1.0) self._ndvi_buf *= np.float32(0.5) # Coloca NDVI01 no último canal if self.channels == 4: raw[3, :, :] = self._ndvi_buf else: raw[4, :, :] = self._ndvi_buf return raw def _preprocess_tensor(self, raw_np: np.ndarray) -> torch.Tensor: """ raw_np: (C,H,W) float32 em 0..1, C=4 ou 5. Retorna imgs_norm: [1,C,H,W] já normalizado por normalize_raw. """ if raw_np.dtype != np.float32: raw_np = raw_np.astype(np.float32) x = torch.from_numpy(raw_np) # [C,H,W] x = x.unsqueeze(0).to(self.device) # [1,C,H,W] x = normalize_raw(x) # z-score por canal (igual treino) return x def _build_preview(self, raw_np, pred_ids, alpha, fast): t0 = time.time() bgr = make_bgr_preview_from_raw(raw_np, rgirb=False, preview_fast=fast) pred_bgr = converter_mask_ids_para_bgr(pred_ids, self.colormap_rgb, self.ignore_id) a = float(np.clip(alpha, 0.0, 1.0)) overlay = (bgr.astype(np.float32) * (1 - a) + pred_bgr.astype(np.float32) * a) overlay = np.clip(overlay, 0, 255).astype(np.uint8) dt = (time.time() - t0) * 1000.0 # ms return bgr, pred_bgr, overlay, dt # ------------------------------------------------- # Inferência # ------------------------------------------------- @torch.no_grad() def infer_ids(self, raw_np: np.ndarray, use_amp: bool | None = None) -> np.ndarray: """ raw_np: (C,H,W) float32 (0..1), C=4 ou 5. Retorna pred_ids [H,W] uint8. """ if use_amp is None: use_amp = self.use_amp imgs_norm = self._preprocess_tensor(raw_np) # [1,C,H,W] ctx = amp.autocast("cuda") if (use_amp and self.device.type == "cuda") else nullcontext() with ctx: out = self.model(pixel_values=imgs_norm) logits = out.logits # [B,C,h,w] logits = torch.nn.functional.interpolate( logits, size=imgs_norm.shape[-2:], mode="bilinear", align_corners=False, ) pred = torch.argmax(logits, dim=1) pred_ids = pred.squeeze(0).cpu().numpy().astype(np.uint8) return pred_ids @torch.inference_mode() def infer_ids_v2(self, raw_np: np.ndarray, use_amp: bool | None = None) -> np.ndarray: """ raw_np: (C,H,W) float32 0..1 Retorna: pred_ids (H,W) uint8 """ if use_amp is None: use_amp = self.use_amp and (self.device.type == "cuda") if raw_np.dtype != np.float32: raw_np = raw_np.astype(np.float32, copy=False) C, H, W = raw_np.shape self.prepare_infer_buffers(C, H, W) # Copia raw_np -> pinned tensor (1,C,H,W) sem criar tensor novo no GPU ainda # (copy_ é rápido) self._cpu_pinned[0].copy_(torch.from_numpy(raw_np), non_blocking=True) # H2D non_blocking x = self._gpu_input x.copy_(self._cpu_pinned, non_blocking=True) # Normalização FIXA (recomendado) if self._norm_mean is not None and self._norm_std is not None: x = (x - self._norm_mean) / self._norm_std else: # fallback (mais lento): normalização por frame mean = x.mean(dim=(2, 3), keepdim=True) std = x.std(dim=(2, 3), keepdim=True).clamp_min(1e-6) x = (x - mean) / std ctx = amp.autocast("cuda") if use_amp else nullcontext() with ctx: out = self.model(pixel_values=x) logits = out.logits # (1,num_classes,h,w) # Upsample só se precisar if logits.shape[-2:] != (H, W): logits = F.interpolate(logits, size=(H, W), mode="bilinear", align_corners=False) pred = torch.argmax(logits, dim=1) # (1,H,W) # puxa pra CPU pred_ids = pred[0].to(torch.uint8).cpu().numpy() return pred_ids @torch.no_grad() def infer_raw(self, raw_np: np.ndarray): """ Pacote completo: - pega entrada RAW (C canais) - faz inferência Retorna: (pred_ids) """ #pred_ids = self.infer_ids(raw_np) pred_ids = self.infer_ids_v2(raw_np) return pred_ids @torch.no_grad() def preview_infer_cached(self, raw_np, pred_ids, alpha=0.5): """ Retorna sempre algo imediatamente: - se preview novo estiver pronto → retorna ele - se ainda estiver processando → retorna último preview válido Retorna: (rgb, pred_rgb, overlay, ts, dt_ms) """ # Se não tem preview ainda, força gerar o primeiro (bloqueante) if self._last_preview is None: rgb, pred_rgb, overlay, dt = self._build_preview(raw_np, pred_ids, alpha, False) self._last_preview = (rgb, pred_rgb, overlay) self._last_preview_ts = time.time() self._last_preview_dt = dt return rgb, pred_rgb, overlay, self._last_preview_ts, dt # Se já tem preview, tenta disparar um novo em background if not self._preview_busy: def worker(): try: rgb, pred_rgb, overlay, dt = self._build_preview(raw_np, pred_ids, alpha, False) with self._preview_lock: self._last_preview = (rgb, pred_rgb, overlay) self._last_preview_ts = time.time() self._last_preview_dt = dt finally: self._preview_busy = False self._preview_busy = True threading.Thread(target=worker, daemon=True).start() # Retorna o último preview imediatamente (cache) rgb, pred_rgb, overlay = self._last_preview return ( rgb, pred_rgb, overlay, self._last_preview_ts, self._last_preview_dt, ) @torch.no_grad() def infer_and_preview(self, raw_np: np.ndarray, alpha=0.5): """ Pacote completo: - pega entrada RAW (C canais) - faz inferência Retorna: (pred_ids) """ #pred_ids = self.infer_ids(raw_np) t0 = time.time() pred_ids = self.infer_ids_v2(raw_np) t1 = time.time() rgb, pred_rgb, overlay, dt = self._build_preview(raw_np, pred_ids, alpha, False) t2 = time.time() t_inf = (t1 - t0) * 1000 t_pvw = (t2 - t1) * 1000 return pred_ids, rgb, pred_rgb, overlay, t_inf, t_pvw # ------------------------------------------------- # Helpers de introspecção # ------------------------------------------------- def get_classes(self): return self.classes def get_colormap(self): return self.colormap_rgb def get_ignore_id(self): return self.ignore_id class RawSegDataset(Dataset): """ Espera duas formas de estrutura: 1) Simples (sem grupos): root/ raws/*.npy|.npz|.raw masks/*.png 2) Com grupos (nosso caso atual): root/ group/ / raws/*.npy|.npz|.raw masks/*.png / ... labelmap.txt fica em dataset/labelmap.txt (passado via labelmap_path). Saída: - image: tensor (C,H,W) com C=4 ou 5 dependendo de 'channels' e 'use_ndvi' - mask : (H,W) long """ def __init__(self, root: str, labelmap_path: Optional[str] = None, raw_exts: Tuple[str, ...] = (".npy", ".npz", ".raw"), max_value: Optional[float] = None, resize_hw: Optional[Tuple[int, int]] = None, require_four_channels: bool = True, raw_hw: Optional[Tuple[int, int]] = None, use_ndvi: bool = False, ndvi_eps: float = 1e-6, channels: int = 4): super().__init__() self.root = root self.raw_exts = raw_exts self.max_value = max_value self.resize_hw = resize_hw self.require_four_channels = require_four_channels self.raw_hw = raw_hw self.use_ndvi = use_ndvi self.ndvi_eps = ndvi_eps self.channels = int(channels) if self.channels not in (4, 5): raise ValueError(f"[dataset] 'channels' inválido: {self.channels} (esperado 4 ou 5)") if self.channels == 5 and not self.use_ndvi: raise ValueError("[dataset] channels=5 sem NDVI não é suportado. " "Use channels=4 sem NDVI ou channels=5 com NDVI.") self.items: List[Tuple[str, str]] = [] # (raw_path, mask_path) dir_group = os.path.join(root, "group") dir_raw_simple = os.path.join(root, "raws") dir_mask_simple = os.path.join(root, "masks") if os.path.isdir(dir_group): # modo "com grupos" for g in sorted(os.listdir(dir_group)): gdir = os.path.join(dir_group, g) if not os.path.isdir(gdir): continue g_raw = os.path.join(gdir, "raws") g_msk = os.path.join(gdir, "masks") if not (os.path.isdir(g_raw) and os.path.isdir(g_msk)): continue raws = [ fn for fn in sorted(os.listdir(g_raw)) if fn.lower().endswith(raw_exts) ] for fn in raws: stem = os.path.splitext(fn)[0] m1 = os.path.join(g_msk, stem + ".png") m2 = os.path.join(g_msk, stem + ".PNG") if os.path.exists(m1): self.items.append((os.path.join(g_raw, fn), m1)) elif os.path.exists(m2): self.items.append((os.path.join(g_raw, fn), m2)) else: # modo "simples" (sem grupos) if not os.path.isdir(dir_raw_simple): raise RuntimeError(f"Não achei pasta raws: {dir_raw_simple}") if not os.path.isdir(dir_mask_simple): raise RuntimeError(f"Não achei pasta masks: {dir_mask_simple}") raws = [ fn for fn in sorted(os.listdir(dir_raw_simple)) if fn.lower().endswith(raw_exts) ] for fn in raws: stem = os.path.splitext(fn)[0] m1 = os.path.join(dir_mask_simple, stem + ".png") m2 = os.path.join(dir_mask_simple, stem + ".PNG") if os.path.exists(m1): self.items.append((os.path.join(dir_raw_simple, fn), m1)) elif os.path.exists(m2): self.items.append((os.path.join(dir_raw_simple, fn), m2)) if not self.items: raise RuntimeError(f"Nenhum par raw/mask encontrado em {root} (nomes base precisam bater).") # classes (labelmap opcional) _, _, classes, ignore_rgb = carregar_labelmap_completo(labelmap_path) ignore_id = _infer_ignore_id(ignore_rgb, default_id=255) self.classes = classes self.class_ids = sorted(classes.keys()) # ex.: [0,1,2] self.ignore_id = ignore_id print(f"[dataset] root={root} items={len(self.items)} channels={self.channels} use_ndvi={self.use_ndvi}") def __len__(self): return len(self.items) def _load_raw(self, path: str) -> np.ndarray: ext = os.path.splitext(path)[1].lower() if ext == ".npz": z = np.load(path) key = list(z.keys())[0] arr = z[key] elif ext == ".npy": arr = np.load(path) elif ext == ".raw": if self.raw_hw is None: raise RuntimeError( f"Para ler .raw preciso de raw_hw=(H,W). " f"Passe raw_hw ao criar o dataset." ) H, W = self.raw_hw size_bytes = os.path.getsize(path) mosa_bytes = H * W # mosaico uint8 raw4_bytes = 4 * H * W * 4 # 4 canais * H * W * 4 bytes (float32) if size_bytes == mosa_bytes: # MODO ANTIGO: mosaico uint8 arr = np.fromfile(path, dtype=np.uint8) raw2d = arr.reshape(H, W) if (H % 2) != 0 or (W % 2) != 0: raise RuntimeError(f"raw_hw deve ser par em H e W p/ mosaico 2x2 (R,G,IR,B), veio H={H}, W={W}") H2, W2 = H // 2, W // 2 r_sub = raw2d[0::2, 0::2] g_sub = raw2d[0::2, 1::2] ir_sub = raw2d[1::2, 0::2] b_sub = raw2d[1::2, 1::2] from PIL import Image def upsample(ch_2d: np.ndarray) -> np.ndarray: im = Image.fromarray(ch_2d) # uint8 im = im.resize((W, H), resample=Image.BILINEAR) return np.array(im) r_full = upsample(r_sub) g_full = upsample(g_sub) ir_full = upsample(ir_sub) b_full = upsample(b_sub) # (H, W, 4) uint8 [R,G,B,IR] arr = np.stack([r_full, g_full, b_full, ir_full], axis=-1).astype(np.uint8) elif size_bytes == raw4_bytes: # NOVO MODO: RAW4 float32 salvo pelo normalize arr_f32 = np.fromfile(path, dtype=np.float32) raw4 = arr_f32.reshape(4, H, W) # (C,H,W) [R,G,IR,B] arr = np.transpose(raw4, (1, 2, 0)) # (H,W,4) [R,G,IR,B] # REORGANIZA PARA CONTRATO INTERNO: [R,G,B,IR] arr = arr[..., [0, 1, 3, 2]] else: raise RuntimeError( f"Tamanho inesperado em {path}: {size_bytes} bytes " f"(esperado {mosa_bytes} ou {raw4_bytes})" ) else: raise RuntimeError(f"Extensão não suportada para RAW: {ext} ({path})") # Até aqui queremos (H,W,4) base if arr.ndim != 3: raise RuntimeError(f"RAW precisa ser 3D, veio {arr.shape} em {path}") if arr.shape[0] == 4 and arr.shape[-1] != 4: arr = np.transpose(arr, (1, 2, 0)) # (H,W,4) if self.require_four_channels and arr.shape[-1] != 4: raise RuntimeError(f"Esperava 4 canais base, veio shape={arr.shape} em {path}") return arr # (H,W,4) def _scale_to_float01(self, raw: np.ndarray) -> np.ndarray: if raw.dtype == np.uint16: mv = float(self.max_value) if self.max_value is not None else 65535.0 elif raw.dtype == np.uint8: mv = float(self.max_value) if self.max_value is not None else 255.0 else: # float já vem normalizado ou não; se passar max_value, aplica mv = float(self.max_value) if self.max_value is not None else None raw_f = raw.astype(np.float32) if mv is not None and mv > 0: raw_f = raw_f / mv # clip pra não explodir raw_f = np.clip(raw_f, 0.0, 1.0) return raw_f def _resize(self, x: np.ndarray, new_hw: Tuple[int, int], is_mask: bool) -> np.ndarray: # resize sem opencv: PIL from PIL import Image h, w = new_hw if is_mask: im = Image.fromarray(x) im = im.resize((w, h), resample=Image.NEAREST) return np.array(im) else: # x (H,W,C) im = Image.fromarray((x * 255.0).astype(np.uint8)) im = im.resize((w, h), resample=Image.BILINEAR) y = np.array(im).astype(np.float32) / 255.0 if y.ndim == 2: y = y[..., None] return y def __getitem__(self, idx: int): raw_path, mask_path = self.items[idx] raw = self._load_raw(raw_path) # (H,W,4) uint8 raw = self._scale_to_float01(raw) # float32 0..1 # Montagem de canais finais if self.use_ndvi: if self.channels == 4: # substitui IR por NDVI01: [R,G,B,NDVI01] raw = self._add_ndvi_replace(raw) elif self.channels == 5: # adiciona NDVI: [R,G,B,IR,NDVI01] raw = self._add_ndvi_append(raw) else: raise RuntimeError(f"[dataset] channels inválido em __getitem__: {self.channels}") else: # modo clássico: [R,G,IR,B] if self.channels != 4: raise RuntimeError("[dataset] channels=5 sem NDVI não suportado em __getitem__.") mask = self._imread_mask_png(mask_path) # (H,W) uint8 if self.resize_hw is not None: raw = self._resize(raw, self.resize_hw, is_mask=False) mask = self._resize(mask, self.resize_hw, is_mask=True) # raw: (H,W,C_final) x = torch.from_numpy(raw).permute(2, 0, 1).contiguous() # (C,H,W) y = torch.from_numpy(mask).long() return {"image": x, "mask": y} def _imread_mask_png(self, path: str) -> np.ndarray: # Leitura PNG sem depender de opencv; usa imageio se tiver, senão PIL. try: import imageio.v2 as imageio m = imageio.imread(path) except Exception: from PIL import Image m = np.array(Image.open(path)) # Se vier RGB, pega um canal if m.ndim == 3: m = m[..., 0] return m.astype(np.uint8) def _add_ndvi_replace(self, raw: np.ndarray) -> np.ndarray: """ raw: (H, W, 4) em float32 [0,1], na ordem base [R, G, B, IR]. Retorna (H, W, 4) na ordem [R, G, B, NDVI01], onde NDVI01 está em [0,1]. """ if raw.shape[-1] != 4: raise RuntimeError(f"Esperava 4 canais pra NDVI (replace), veio {raw.shape}") r = raw[..., 0] g = raw[..., 1] # não entra no NDVI, mas mantemos b = raw[..., 2] ir = raw[..., 3] num = ir - r den = ir + r + self.ndvi_eps ndvi = num / den ndvi = np.clip(ndvi, -1.0, 1.0) ndvi01 = (ndvi + 1.0) * 0.5 out = np.stack([r, g, b, ndvi01], axis=-1).astype(np.float32) return out def _add_ndvi_append(self, raw: np.ndarray) -> np.ndarray: """ raw: (H, W, 4) em float32 [0,1], na ordem base [R, G, B, IR]. Retorna (H, W, 5) na ordem [R, G, B, IR, NDVI01], onde NDVI01 está em [0,1]. """ if raw.shape[-1] != 4: raise RuntimeError(f"Esperava 4 canais pra NDVI (append), veio {raw.shape}") r = raw[..., 0] g = raw[..., 1] b = raw[..., 2] ir = raw[..., 3] num = ir - r den = ir + r + self.ndvi_eps ndvi = num / den ndvi = np.clip(ndvi, -1.0, 1.0) ndvi01 = (ndvi + 1.0) * 0.5 out = np.stack([r, g, b, ir, ndvi01], axis=-1).astype(np.float32) return out # ---------------------------- # Normalização RAW (4 ou 5 canais) # ---------------------------- def normalize_raw(imgs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor: """ imgs: (B,C,H,W) em 0..1, com C=4 ou 5. Faz normalização por canal usando mean/std do batch (rápido e robusto pra raw). """ mean = imgs.mean(dim=(0, 2, 3), keepdim=True) std = imgs.std(dim=(0, 2, 3), keepdim=True).clamp_min(eps) return (imgs - mean) / std def make_bgr_preview_from_raw(raw_np: np.ndarray, rgirb: bool, preview_fast: bool, preview_scale: int = 2) -> np.ndarray: """ raw_np: (C,H,W) float32 em 0..1, C=4 ou 5. Layout esperado: - channels=4, use_ndvi=False -> [R,G,B,IR] - channels=4, use_ndvi=True -> [R,G,B,NDVI] - channels=5, use_ndvi=True -> [R,G,B,IR,NDVI] Usamos sempre R,G,B para o preview. """ assert raw_np.ndim == 3, "raw_np deve ser (C,H,W)" C = raw_np.shape[0] if C not in (4, 5): raise RuntimeError(f"Esperado C=4 ou 5, veio {C} em make_rgb_preview_from_raw") r = raw_np[0] g = raw_np[1] b = raw_np[3] if rgirb else raw_np[2] # ------------------------- # Modo PREVIEW RÁPIDO # ------------------------- if preview_fast: # Downscale simples (barato) if preview_scale > 1: r = r[::preview_scale, ::preview_scale] g = g[::preview_scale, ::preview_scale] b = b[::preview_scale, ::preview_scale] bgr = np.stack([b, g, r], axis=0) # Gamma simples só pra não ficar "lavado" bgr = np.power(bgr, 1/1.8) # Converte direto pra uint8 bgr8 = (bgr * 255.0).clip(0, 255).astype(np.uint8) bgr_hwc = np.transpose(bgr8, (1, 2, 0)) # RGB bgr = bgr_hwc.copy() return bgr bgr = np.stack([b, g, r], axis=0).astype(np.float32) # (3,H,W) # Auto white balance simples (referenciado no G) means = bgr.mean(axis=(1, 2)) # (3,) ref = means[1] # canal G gains = ref / (means + 1e-6) gains = np.clip(gains, 0.5, 2.0) bgr_wb = bgr * gains[:, None, None] # Força do WB alpha = 1.0 # teste 0.4..0.8 if alpha >= 0.999: bgr = bgr_wb else: bgr_orig = bgr.copy() bgr = (1 - alpha) * bgr_orig + alpha * bgr_wb bgr = np.clip(bgr, 0.0, 1.0) # Gamma mais suave bgr = np.power(bgr, 1/1.9, out=bgr) # Para OpenCV bgr8 = (bgr * 255.0).clip(0, 255).astype(np.uint8) bgr_hwc = np.transpose(bgr8, (1, 2, 0)) # (H,W,3) RGB bgr = bgr_hwc.copy() return bgr # ---------------------------- # Patch: SegFormer para N canais # ---------------------------- def patch_segformer_input_channels(model: SegformerForSemanticSegmentation, in_ch: int = 4): """ Troca o primeiro conv do encoder (patch embedding stage 0) pra aceitar in_ch. Inicializa canais extras copiando a média dos 3 canais originais, se existirem. """ enc = model.segformer.encoder proj = enc.patch_embeddings[0].proj # Conv2d(in=3,...) if proj.in_channels == in_ch: return old_w = proj.weight.data old_b = proj.bias.data if proj.bias is not None else None new_proj = nn.Conv2d( in_channels=in_ch, out_channels=proj.out_channels, kernel_size=proj.kernel_size, stride=proj.stride, padding=proj.padding, dilation=proj.dilation, groups=proj.groups, bias=(proj.bias is not None), padding_mode=proj.padding_mode, ).to(proj.weight.device) with torch.no_grad(): if old_w.shape[1] == 3 and in_ch >= 3: # copia os 3 canais originais new_proj.weight[:, :3, :, :].copy_(old_w) if in_ch > 3: # canais extras inicializados como média dos 3 extra = old_w.mean(dim=1, keepdim=True) # (out,1,kh,kw) for c in range(3, in_ch): new_proj.weight[:, c:c+1, :, :].copy_(extra) else: # fallback genérico nn.init.kaiming_normal_(new_proj.weight, mode="fan_out", nonlinearity="relu") if old_b is not None: new_proj.bias.copy_(old_b) enc.patch_embeddings[0].proj = new_proj def _infer_ignore_id(ignore_rgb, default_id=255): import numpy as _np if isinstance(ignore_rgb, (list, tuple)): if len(ignore_rgb) == 1 and isinstance(ignore_rgb[0], (int, _np.integer)): return int(ignore_rgb[0]) if len(ignore_rgb) == 3: return default_id if isinstance(ignore_rgb, (int, _np.integer)): return int(ignore_rgb) return default_id