ajustes treinamento raw

This commit is contained in:
Diego Freitas 2026-04-13 12:44:19 -03:00
parent 3a40df8492
commit aa17e0b453
10 changed files with 534 additions and 103 deletions

View File

@ -196,7 +196,11 @@ def main():
gain_a = dbg.get("gain_a", None) gain_a = dbg.get("gain_a", None)
gain_d = dbg.get("gain_d", None) gain_d = dbg.get("gain_d", None)
bgr = make_bgr_preview_from_raw(raw4_base, rgirb=True, preview_fast=upscale > 0, preview_scale=upscale) apply_ir_comp = True
ir_k_r = 0.4
ir_k_g = 0.1
ir_k_b = 0.5
bgr = make_bgr_preview_from_raw(raw4_base, rgirb=True, preview_fast=upscale > 0, preview_scale=upscale, apply_ir_comp=apply_ir_comp, ir_k_r=ir_k_r, ir_k_g=ir_k_g, ir_k_b=ir_k_b)
# FPS # FPS
frames += 1 frames += 1
@ -243,6 +247,10 @@ def main():
"exp_raw": int(exp_raw) if exp_raw is not None else None, "exp_raw": int(exp_raw) if exp_raw is not None else None,
"gain_a": int(gain_a) if gain_a is not None else None, "gain_a": int(gain_a) if gain_a is not None else None,
"gain_d": int(gain_d) if gain_d is not None else None, "gain_d": int(gain_d) if gain_d is not None else None,
"apply_ir_comp": apply_ir_comp,
"ir_k_r": ir_k_r,
"ir_k_g": ir_k_g,
"ir_k_b": ir_k_b,
"ae_dbg": { "ae_dbg": {
k: (float(v) if isinstance(v, (int, float, np.floating)) else v) k: (float(v) if isinstance(v, (int, float, np.floating)) else v)
for k, v in ae_dbg.items() for k, v in ae_dbg.items()

View File

@ -43,8 +43,8 @@ except Exception as e:
raise raise
RGB_SUFFIXES = ["_rgb", "_RGB"] RGB_SUFFIXES = ["_rgb", "_RGB", "_Rgb"]
SEG_SUFFIXES = ["_segmentacao", "_seg", "_SEG", "_segment"] # allow a few variations SEG_SUFFIXES = ["_segmentacao", "_seg", "_SEG", "_segment", "_Segmentacao"] # allow a few variations
EXTS = [".jpeg", ".jpg", ".png"] EXTS = [".jpeg", ".jpg", ".png"]

View File

@ -19,8 +19,8 @@ PASTA_FINAL_IMAGES = os.path.join(MODELO, "dataset", "original", "images") #
PASTA_FINAL_MASKS = os.path.join(MODELO, "dataset", "original", "masks") # dataset final (máscaras) PASTA_FINAL_MASKS = os.path.join(MODELO, "dataset", "original", "masks") # dataset final (máscaras)
# Classe-alvo (RGB) para preencher a máscara sólida # Classe-alvo (RGB) para preencher a máscara sólida
COR_CLASSE_RGB = (128, 0, 0) # (R, G, B) COR_CLASSE_RGB = (0, 128, 0) # (R, G, B)
COPIAR_IMAGENS = True COPIAR_IMAGENS = False
# Controle # Controle
EXT_IMAGENS = (".jpg", ".jpeg", ".png") EXT_IMAGENS = (".jpg", ".jpeg", ".png")

View File

@ -274,8 +274,10 @@ def _load_raw(path: str, raw_hw) -> np.ndarray:
g_full = upsample(g_sub) g_full = upsample(g_sub)
ir_full = upsample(ir_sub) ir_full = upsample(ir_sub)
b_full = upsample(b_sub) b_full = upsample(b_sub)
# (H, W, 4) uint8 [R,G,B,IR] # (H, W, 4) uint8 [R,G,IR,B]
arr = np.stack([r_full, g_full, b_full, ir_full], axis=-1).astype(np.uint8) arr = np.stack([r_full, g_full, ir_full, b_full], axis=-1).astype(np.uint8)
# REORGANIZA PARA CONTRATO INTERNO [R,G,B,IR]
arr = arr[..., [0, 1, 3, 2]]
elif size_bytes == raw4_bytes: elif size_bytes == raw4_bytes:
# NOVO MODO: RAW4 float32 salvo pelo normalize # NOVO MODO: RAW4 float32 salvo pelo normalize
arr_f32 = np.fromfile(path, dtype=np.float32) arr_f32 = np.fromfile(path, dtype=np.float32)
@ -317,7 +319,14 @@ def _scale_to_float01(raw: np.ndarray, max_value = None) -> np.ndarray:
def save_raw_any(path, arr): def save_raw_any(path, arr):
arr = np.asarray(arr) arr = np.asarray(arr)
arr.tofile(path)
if arr.ndim == 3 and arr.shape[-1] == 4:
# arr está no contrato interno [R,G,B,IR]
arr_cam = internal_rgbir_to_camera_raw4(arr) # [R,G,IR,B]
arr_cam = np.transpose(arr_cam, (2, 0, 1)) # (4,H,W)
arr_cam.astype(np.float32).tofile(path)
else:
arr.tofile(path)
# =============================== # ===============================
@ -424,6 +433,7 @@ def augment_sample(
msk2_path=None, msk2_path=None,
msk2_out_dir=None, msk2_out_dir=None,
raw_out_dir=None, raw_out_dir=None,
aug_suffix="aug",
): ):
""" """
Faz a augmentação a partir de UM RAW4 + máscara (e opcionalmente máscara2). Faz a augmentação a partir de UM RAW4 + máscara (e opcionalmente máscara2).
@ -480,17 +490,17 @@ def augment_sample(
img_b = raw_to_preview_rgb(raw_b) img_b = raw_to_preview_rgb(raw_b)
# 7) Salva cópias # 7) Salva cópias
out_img = os.path.join(img_out_dir, f"{base}_aug_{i:02d}.png") # preview em PNG out_img = os.path.join(img_out_dir, f"{base}_{aug_suffix}_{i:02d}.png")
out_msk = os.path.join(msk_out_dir, f"{base}_aug_{i:02d}{msk_ext}") out_msk = os.path.join(msk_out_dir, f"{base}_{aug_suffix}_{i:02d}{msk_ext}")
save_rgb(out_img, img_b) save_rgb(out_img, img_b)
save_rgb(out_msk, msk_g) save_rgb(out_msk, msk_g)
if msk2_g is not None and msk2_out_dir: if msk2_g is not None and msk2_out_dir:
out_msk2 = os.path.join(msk2_out_dir, f"{base}_aug_{i:02d}{msk2_ext}") out_msk2 = os.path.join(msk2_out_dir, f"{base}_{aug_suffix}_{i:02d}{msk2_ext}")
save_rgb(out_msk2, msk2_g) save_rgb(out_msk2, msk2_g)
if raw_b is not None and raw_out_dir: if raw_b is not None and raw_out_dir:
out_raw = os.path.join(raw_out_dir, f"{base}_aug_{i:02d}{raw_ext}") out_raw = os.path.join(raw_out_dir, f"{base}_{aug_suffix}_{i:02d}{raw_ext}")
save_raw_any(out_raw, raw_b) save_raw_any(out_raw, raw_b)
gen += 1 gen += 1
@ -502,7 +512,7 @@ def augment_sample(
# Processamento por grupo # Processamento por grupo
# =============================== # ===============================
def process_group(group_name, copies): def process_group(group_name, copies, limit=None, seed=42, aug_suffix="aug"):
""" """
Processa um grupo único usando APENAS: Processa um grupo único usando APENAS:
- group/<g>/raws - group/<g>/raws
@ -525,6 +535,15 @@ def process_group(group_name, copies):
use_masks2 = USE_MASKS2 and os.path.isdir(msk2_dir) use_masks2 = USE_MASKS2 and os.path.isdir(msk2_dir)
raw_map = map_raws_by_base(raw_dir) raw_map = map_raws_by_base(raw_dir)
items = sorted(raw_map.items())
if limit is not None and limit > 0:
rng = np.random.default_rng(seed)
if limit < len(items):
idx = rng.choice(len(items), size=limit, replace=False)
idx = sorted(idx.tolist())
items = [items[i] for i in idx]
msk_map = map_by_base_priorizando_png(msk_dir, MSK_EXTS) msk_map = map_by_base_priorizando_png(msk_dir, MSK_EXTS)
msk2_map = map_by_base_priorizando_png(msk2_dir, MSK2_EXTS) if use_masks2 else {} msk2_map = map_by_base_priorizando_png(msk2_dir, MSK2_EXTS) if use_masks2 else {}
@ -536,7 +555,7 @@ def process_group(group_name, copies):
) )
count = 0 count = 0
for base, raw_file in sorted(raw_map.items()): for base, raw_file in items:
msk_file = msk_map.get(base) msk_file = msk_map.get(base)
if not msk_file: if not msk_file:
print(f"[WARN] [{group_name}] Máscara não encontrada para RAW {base}, pulando.") print(f"[WARN] [{group_name}] Máscara não encontrada para RAW {base}, pulando.")
@ -556,6 +575,7 @@ def process_group(group_name, copies):
msk2_path=msk2_file, msk2_path=msk2_file,
msk2_out_dir=msk2_out_dir, msk2_out_dir=msk2_out_dir,
raw_out_dir=raw_out_dir, raw_out_dir=raw_out_dir,
aug_suffix=aug_suffix,
) )
except Exception as e: except Exception as e:
print(f"[ERRO] [{group_name}] {base}: {e}") print(f"[ERRO] [{group_name}] {base}: {e}")
@ -633,12 +653,23 @@ def process_legacy(copies):
print(f"[OK] Legacy → {count} amostras geradas.") print(f"[OK] Legacy → {count} amostras geradas.")
return count return count
def internal_rgbir_to_camera_raw4(arr: np.ndarray) -> np.ndarray:
"""
Converte do contrato interno (H,W,4) [R,G,B,IR]
para o contrato bruto da câmera (H,W,4) [R,G,IR,B].
"""
if arr is None:
return None
if arr.ndim != 3 or arr.shape[-1] != 4:
raise RuntimeError(f"Esperado array (H,W,4), veio {arr.shape}")
return arr[..., [0, 1, 3, 2]]
# =============================== # ===============================
# main # main
# =============================== # ===============================
def main(copies=5, groups_csv=None): def main(copies=5, groups_csv=None, limit=None, seed=42, suffix="aug"):
total = 0 total = 0
if os.path.isdir(ORIG_GROUP_ROOT): if os.path.isdir(ORIG_GROUP_ROOT):
grupos = list_groups(ORIG_GROUP_ROOT) grupos = list_groups(ORIG_GROUP_ROOT)
@ -653,7 +684,7 @@ def main(copies=5, groups_csv=None):
else: else:
print(f"Grupos encontrados: {', '.join(grupos)}") print(f"Grupos encontrados: {', '.join(grupos)}")
for g in grupos: for g in grupos:
total += process_group(g, copies) total += process_group(g, copies, limit=limit, seed=seed, aug_suffix=suffix)
else: else:
total += process_legacy(copies) total += process_legacy(copies)
@ -664,5 +695,14 @@ if __name__ == "__main__":
ap = argparse.ArgumentParser(description="Augmentação por grupos usando RAW4 (preview gerado do RAW).") ap = argparse.ArgumentParser(description="Augmentação por grupos usando RAW4 (preview gerado do RAW).")
ap.add_argument("--copies", type=int, default=5, help="Número de cópias augmentadas por imagem (default=5).") ap.add_argument("--copies", type=int, default=5, help="Número de cópias augmentadas por imagem (default=5).")
ap.add_argument("--groups", type=str, default=None, help="Lista de grupos separados por vírgula (ex: chao,erva_cana).") ap.add_argument("--groups", type=str, default=None, help="Lista de grupos separados por vírgula (ex: chao,erva_cana).")
ap.add_argument("--limit", type=int, default=None, help="Quantidade máxima de imagens originais do grupo a augmentar.")
ap.add_argument("--seed", type=int, default=42, help="Seed para seleção reproduzível quando usar --limit.")
ap.add_argument("--suffix", type=str, default="aug", help="Sufixo usado no nome dos arquivos gerados.")
args = ap.parse_args() args = ap.parse_args()
main(copies=args.copies, groups_csv=args.groups) main(
copies=args.copies,
groups_csv=args.groups,
limit=args.limit,
seed=args.seed,
suffix=args.suffix,
)

View File

@ -42,8 +42,8 @@ RAW_EXT = ".raw"
# Regex para identificar famílias # Regex para identificar famílias
RE_ORIGINAL_PREFIX = re.compile(r'^original_(.+)$', re.IGNORECASE) RE_ORIGINAL_PREFIX = re.compile(r'^original_(.+)$', re.IGNORECASE)
RE_AUGMENTED_FAMILY = re.compile(r'^augmented_(.+?)(?:_aug_\d+)?$', re.IGNORECASE) RE_AUGMENTED_FAMILY = re.compile(r'^augmented_(.+?)(?:_aug[a-zA-Z0-9]*_\d+)?$', re.IGNORECASE)
RE_AUG_SUFFIX = re.compile(r'_aug_\d+$', re.IGNORECASE) RE_AUG_SUFFIX = re.compile(r'_aug[a-zA-Z0-9]*_\d+$', re.IGNORECASE)
def garantir(p): def garantir(p):

View File

@ -16,7 +16,7 @@ from torch.utils.data import DataLoader, Dataset
from torch.amp import autocast, GradScaler from torch.amp import autocast, GradScaler
import torch.nn.functional as F import torch.nn.functional as F
from raw_segformer_service import (RawSegDataset, normalize_raw, patch_segformer_input_channels, build_raw_segformer_model) from raw_segformer_service import (RawSegDataset, normalize_raw, patch_segformer_input_channels, build_raw_segformer_model, build_dual_branch_segformer_model)
# ---------------------------- # ----------------------------
@ -327,6 +327,7 @@ def main():
parser.add_argument("--resize_w", type=int, default=None) parser.add_argument("--resize_w", type=int, default=None)
parser.add_argument("--norm_stats", type=str, default=None, help="Caminho para JSON com mean/std por canal (ex: norm_stats.json).") parser.add_argument("--norm_stats", type=str, default=None, help="Caminho para JSON com mean/std por canal (ex: norm_stats.json).")
parser.add_argument("--fusion_mode", type=str, default=None)
args = parser.parse_args() args = parser.parse_args()
set_seed(args.seed) set_seed(args.seed)
@ -348,7 +349,18 @@ def main():
if args.main_class is not None: if args.main_class is not None:
MAIN_CLASS_NAME = args.main_class.lower() MAIN_CLASS_NAME = args.main_class.lower()
save_path = os.path.join(MODELO, "backup", config["modelo"], MODEL_NAME, f"raw{CHANNELS}") FUSION_MODE = config.get("fusion_mode", "stacked")
if args.fusion_mode is not None:
FUSION_MODE = args.fusion_mode
if FUSION_MODE == "dual_branch" and CHANNELS not in (4, 5):
raise ValueError("dual_branch requer channels=4 ou 5")
stats_source_tag = config.get("stats_source_tag", "stacked_raw4")
experiment_tag = f"{FUSION_MODE}_raw{CHANNELS}"
save_path = os.path.join(MODELO, "backup", config["modelo"], MODEL_NAME, experiment_tag)
#save_path = os.path.join(MODELO, "backup", config["modelo"], MODEL_NAME, f"raw{CHANNELS}")
dataset_path = os.path.join(MODELO, "dataset") dataset_path = os.path.join(MODELO, "dataset")
labelmap_path = os.path.join(dataset_path, "labelmap.txt") labelmap_path = os.path.join(dataset_path, "labelmap.txt")
os.makedirs(save_path, exist_ok=True) os.makedirs(save_path, exist_ok=True)
@ -365,7 +377,8 @@ def main():
# Caminho padrão: dentro do dataset, nome do arquivo de stats # Caminho padrão: dentro do dataset, nome do arquivo de stats
# (ajusta aqui pro nome que você realmente usou: norm_stats.json, por ex.) # (ajusta aqui pro nome que você realmente usou: norm_stats.json, por ex.)
norm_stats_path = os.path.join(save_path, "norm_stats.json") #norm_stats_path = os.path.join(save_path, "norm_stats.json")
norm_stats_path = os.path.join(MODELO, "backup", config["modelo"], MODEL_NAME, stats_source_tag, "norm_stats.json")
if args.norm_stats is not None: if args.norm_stats is not None:
norm_stats_path = args.norm_stats norm_stats_path = args.norm_stats
@ -411,7 +424,19 @@ def main():
mean_rgbi = np.array(mean_rgbi, dtype=np.float32) mean_rgbi = np.array(mean_rgbi, dtype=np.float32)
std_rgbi = np.array(std_rgbi, dtype=np.float32) std_rgbi = np.array(std_rgbi, dtype=np.float32)
if CHANNELS == 4 and not USE_NDVI: if CHANNELS == 3 and not USE_NDVI:
mean3 = mean_rgbi[:3]
std3 = std_rgbi[:3]
mean_t = torch.tensor(mean3, dtype=torch.float32, device=device).view(1, 3, 1, 1)
std_t = torch.tensor(std3, dtype=torch.float32, device=device).view(1, 3, 1, 1)
def normalizer(x: torch.Tensor) -> torch.Tensor:
return (x - mean_t) / std_t
print("[NORM] Normalização fixa por canal ativada para [R,G,B].")
elif CHANNELS == 4 and not USE_NDVI:
# [R,G,B,IR] # [R,G,B,IR]
mean_t = torch.tensor(mean_rgbi, dtype=torch.float32, device=device).view(1, 4, 1, 1) mean_t = torch.tensor(mean_rgbi, dtype=torch.float32, device=device).view(1, 4, 1, 1)
std_t = torch.tensor(std_rgbi, dtype=torch.float32, device=device).view(1, 4, 1, 1) std_t = torch.tensor(std_rgbi, dtype=torch.float32, device=device).view(1, 4, 1, 1)
@ -512,17 +537,28 @@ def main():
drop_last=False, drop_last=False,
) )
model = build_raw_segformer_model( if FUSION_MODE == "stacked":
num_classes=num_classes, model = build_raw_segformer_model(
channels=CHANNELS, num_classes=num_classes,
backbone=BACKBONE, channels=CHANNELS,
device=device, backbone=BACKBONE,
ckpt_path=None, # treino começa do backbone ADE device=device,
strict=False, # pode deixar False se quiser mais flexibilidade ckpt_path=None,
) strict=False,
)
patch_segformer_input_channels(model, in_ch=CHANNELS)
# patch para N canais de entrada (4 ou 5) elif FUSION_MODE == "dual_branch":
patch_segformer_input_channels(model, in_ch=CHANNELS) model = build_dual_branch_segformer_model(
num_classes=num_classes,
channels=CHANNELS,
backbone=BACKBONE,
device=device,
ckpt_path=None,
strict=False,
)
else:
raise ValueError(f"fusion_mode inválido: {FUSION_MODE}")
if args.grad_ckpt: if args.grad_ckpt:
try: try:

View File

@ -232,6 +232,7 @@ def main():
USE_NDVI = bool(config.get("use_ndvi", False)) USE_NDVI = bool(config.get("use_ndvi", False))
CHANNELS = int(config.get("channels", 4)) CHANNELS = int(config.get("channels", 4))
FUSION_MODE = config.get("fusion_mode", "stacked")
RESOLUCAO = config["resolucao"] RESOLUCAO = config["resolucao"]
W, H = RESOLUCAO[0], RESOLUCAO[1] W, H = RESOLUCAO[0], RESOLUCAO[1]
if args.resize_h is not None: if args.resize_h is not None:
@ -251,7 +252,9 @@ def main():
norm_mean = None norm_mean = None
norm_std = None norm_std = None
norm_stats_path = os.path.join(MODELO, "backup", modelo_folder, MODEL_NAME, f"raw{CHANNELS}", "norm_stats.json") experiment_tag = f"{FUSION_MODE}_raw{CHANNELS}"
norm_stats_path = os.path.join(MODELO, "backup", modelo_folder, MODEL_NAME, experiment_tag, "norm_stats.json")
if args.norm_stats is not None: if args.norm_stats is not None:
norm_stats_path = args.norm_stats norm_stats_path = args.norm_stats
if os.path.isfile(norm_stats_path): if os.path.isfile(norm_stats_path):
@ -279,7 +282,10 @@ def main():
s_IR = stats_std[idx_by_name["IR"]] s_IR = stats_std[idx_by_name["IR"]]
s_B = stats_std[idx_by_name["B"]] s_B = stats_std[idx_by_name["B"]]
if CHANNELS == 4 and not USE_NDVI: if CHANNELS == 3 and not USE_NDVI:
norm_mean = [m_R, m_G, m_B]
norm_std = [s_R, s_G, s_B]
elif CHANNELS == 4 and not USE_NDVI:
norm_mean = [m_R, m_G, m_B, m_IR] norm_mean = [m_R, m_G, m_B, m_IR]
norm_std = [s_R, s_G, s_B, s_IR] norm_std = [s_R, s_G, s_B, s_IR]
elif CHANNELS == 4 and USE_NDVI: elif CHANNELS == 4 and USE_NDVI:
@ -301,7 +307,7 @@ def main():
if args.ckpt is not None: if args.ckpt is not None:
ckpt_path = args.ckpt ckpt_path = args.ckpt
else: else:
save_path = os.path.join(MODELO, "backup", modelo_folder, MODEL_NAME, f"raw{CHANNELS}") save_path = os.path.join(MODELO, "backup", modelo_folder, MODEL_NAME, experiment_tag)
ckpt_path = os.path.join(save_path, "best_miou.pt") ckpt_path = os.path.join(save_path, "best_miou.pt")
if not os.path.isfile(ckpt_path): if not os.path.isfile(ckpt_path):

View File

@ -1,7 +1,7 @@
{ {
"camera": "gal5000", "camera": "gal5000",
"modelo": "segformer_b1", "modelo": "segformer_b1",
"model_name": "pulv_ref", "model_name": "pulv_new",
"dual_head": false, "dual_head": false,
"main_class_name": "cana", "main_class_name": "cana",
"es_classes": "", "es_classes": "",
@ -11,7 +11,9 @@
"roi_inicio": 0.0, "roi_inicio": 0.0,
"roi_tamanho": 1.0, "roi_tamanho": 1.0,
"shaves": 3, "shaves": 3,
"channels": 5, "channels": 4,
"use_ndvi": true, "use_ndvi": false,
"backbone": "nvidia/segformer-b1-finetuned-ade-512-512" "backbone": "nvidia/segformer-b1-finetuned-ade-512-512",
"fusion_mode": "stacked",
"stats_source_tag": "stacked_raw4"
} }

View File

@ -1,7 +1,7 @@
{ {
"camera": "oak-d", "camera": "oak-d",
"modelo": "segformer_b0", "modelo": "segformer_b0",
"model_name": "nav", "model_name": "nav_big",
"dual_head": false, "dual_head": false,
"main_class_name": "navegavel", "main_class_name": "navegavel",
"es_classes": "", "es_classes": "",
@ -12,6 +12,6 @@
"roi_tamanho": 1.0, "roi_tamanho": 1.0,
"shaves": 3, "shaves": 3,
"channels": 3, "channels": 3,
"use_ndvi": true, "use_ndvi": false,
"backbone": "nvidia/segformer-b0-finetuned-ade-512-512" "backbone": "nvidia/segformer-b0-finetuned-ade-512-512"
} }

View File

@ -55,6 +55,282 @@ def build_raw_segformer_model(
model.to(device) model.to(device)
return model return model
def build_dual_branch_segformer_model(
num_classes: int,
channels: int,
backbone: str = "nvidia/segformer-b1-finetuned-ade-512-512",
device: torch.device | None = None,
ckpt_path: str | None = None,
strict: bool = True,
) -> nn.Module:
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if channels not in (4, 5):
raise ValueError(f"Dual branch espera channels=4 ou 5, veio {channels}")
spec_channels = channels - 3 # 1 (IR) ou 2 (IR+NDVI)
model = DualBranchSegformerV2(
num_classes=num_classes,
backbone=backbone,
spec_channels=spec_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 SpectralBranch(nn.Module):
def __init__(self, in_ch: int, out_ch: int = 256):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(in_ch, 32, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.Conv2d(128, out_ch, kernel_size=1, stride=1, padding=0, bias=False),
nn.BatchNorm2d(out_ch),
nn.ReLU(inplace=True),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
class DualBranchSegformer(nn.Module):
"""
Entrada esperada:
- 4 canais: [R,G,B,IR]
- 5 canais: [R,G,B,IR,NDVI]
"""
def __init__(
self,
num_classes: int,
backbone: str = "nvidia/segformer-b1-finetuned-ade-512-512",
spec_channels: int = 1,
):
super().__init__()
self.rgb_model = SegformerForSemanticSegmentation.from_pretrained(
backbone,
num_labels=num_classes,
ignore_mismatched_sizes=True,
use_safetensors=True,
)
patch_segformer_input_channels(self.rgb_model, in_ch=3)
self.spec_branch = nn.Sequential(
nn.Conv2d(spec_channels, 32, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.Conv2d(128, num_classes, kernel_size=1, stride=1, padding=0, bias=True),
)
self.fuse = nn.Sequential(
nn.Conv2d(num_classes * 2, num_classes, kernel_size=1, bias=True),
)
def forward(self, pixel_values: torch.Tensor):
rgb = pixel_values[:, :3, :, :]
spec = pixel_values[:, 3:, :, :]
rgb_out = self.rgb_model(pixel_values=rgb)
rgb_logits = rgb_out.logits # (B,num_classes,h,w)
spec_logits = self.spec_branch(spec)
if spec_logits.shape[-2:] != rgb_logits.shape[-2:]:
spec_logits = F.interpolate(
spec_logits,
size=rgb_logits.shape[-2:],
mode="bilinear",
align_corners=False,
)
fused = torch.cat([rgb_logits, spec_logits], dim=1)
logits = self.fuse(fused)
return type("DualBranchOutput", (), {"logits": logits})
class ResidualBlock(nn.Module):
def __init__(self, ch: int):
super().__init__()
self.conv1 = nn.Conv2d(ch, ch, 3, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(ch)
self.act = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(ch, ch, 3, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(ch)
def forward(self, x):
identity = x
x = self.act(self.bn1(self.conv1(x)))
x = self.bn2(self.conv2(x))
x = self.act(x + identity)
return x
class SpectralEncoderV2(nn.Module):
def __init__(self, in_ch: int, out_ch: int = 256):
super().__init__()
self.stem = nn.Sequential(
nn.Conv2d(in_ch, 32, 3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
)
self.stage1 = nn.Sequential(
ResidualBlock(32),
nn.Conv2d(32, 64, 3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
)
self.stage2 = nn.Sequential(
ResidualBlock(64),
nn.Conv2d(64, 128, 3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
)
self.stage3 = nn.Sequential(
ResidualBlock(128),
nn.Conv2d(128, out_ch, 1, bias=False),
nn.BatchNorm2d(out_ch),
nn.ReLU(inplace=True),
)
def forward(self, x):
x = self.stem(x)
x = self.stage1(x)
x = self.stage2(x)
x = self.stage3(x)
return x
class FeatureFusionBlock(nn.Module):
def __init__(self, ch_rgb: int, ch_spec: int, ch_fused: int):
super().__init__()
self.rgb_proj = nn.Conv2d(ch_rgb, ch_fused, 1, bias=False)
self.spec_proj = nn.Conv2d(ch_spec, ch_fused, 1, bias=False)
self.fuse = nn.Sequential(
nn.Conv2d(ch_fused * 2, ch_fused, 3, padding=1, bias=False),
nn.BatchNorm2d(ch_fused),
nn.ReLU(inplace=True),
nn.Conv2d(ch_fused, ch_fused, 3, padding=1, bias=False),
nn.BatchNorm2d(ch_fused),
nn.ReLU(inplace=True),
)
self.gate = nn.Sequential(
nn.Conv2d(ch_fused, ch_fused, 1),
nn.Sigmoid(),
)
def forward(self, rgb_feat, spec_feat):
rgb_feat = self.rgb_proj(rgb_feat)
spec_feat = self.spec_proj(spec_feat)
if spec_feat.shape[-2:] != rgb_feat.shape[-2:]:
spec_feat = F.interpolate(
spec_feat,
size=rgb_feat.shape[-2:],
mode="bilinear",
align_corners=False,
)
fused = torch.cat([rgb_feat, spec_feat], dim=1)
fused = self.fuse(fused)
gate = self.gate(fused)
fused = fused * gate
return fused
class DualBranchSegformerV2(nn.Module):
def __init__(
self,
num_classes: int,
backbone: str = "nvidia/segformer-b1-finetuned-ade-512-512",
spec_channels: int = 1,
fused_channels: int = 256,
):
super().__init__()
self.rgb_model = SegformerForSemanticSegmentation.from_pretrained(
backbone,
num_labels=num_classes,
ignore_mismatched_sizes=True,
use_safetensors=True,
)
patch_segformer_input_channels(self.rgb_model, in_ch=3)
self.spec_branch = SpectralEncoderV2(
in_ch=spec_channels,
out_ch=fused_channels,
)
self.fusion = FeatureFusionBlock(
ch_rgb=self.rgb_model.config.hidden_sizes[-1],
ch_spec=fused_channels,
ch_fused=fused_channels,
)
self.classifier = nn.Sequential(
nn.Conv2d(fused_channels, fused_channels, 3, padding=1, bias=False),
nn.BatchNorm2d(fused_channels),
nn.ReLU(inplace=True),
nn.Conv2d(fused_channels, num_classes, 1, bias=True),
)
def forward(self, pixel_values: torch.Tensor):
rgb = pixel_values[:, :3, :, :]
spec = pixel_values[:, 3:, :, :]
rgb_out = self.rgb_model(
pixel_values=rgb,
output_hidden_states=True,
)
rgb_logits = rgb_out.logits
rgb_feat = rgb_out.hidden_states[-1]
spec_feat = self.spec_branch(spec)
fused_feat = self.fusion(rgb_feat, spec_feat)
logits = self.classifier(fused_feat)
if logits.shape[-2:] != rgb_logits.shape[-2:]:
logits = F.interpolate(
logits,
size=rgb_logits.shape[-2:],
mode="bilinear",
align_corners=False,
)
return type("DualBranchOutput", (), {"logits": logits})
class RawSegformerService: class RawSegformerService:
""" """
Service para trabalhar com o modelo SegFormer RAW (4 ou 5 canais). Service para trabalhar com o modelo SegFormer RAW (4 ou 5 canais).
@ -116,6 +392,7 @@ class RawSegformerService:
self.modelo_folder = cfg["modelo"] # ex: "dif" self.modelo_folder = cfg["modelo"] # ex: "dif"
self.res_w, self.res_h = cfg["resolucao"] # [W,H] self.res_w, self.res_h = cfg["resolucao"] # [W,H]
self.backbone = cfg["backbone"] self.backbone = cfg["backbone"]
self.fusion_mode = cfg.get("fusion_mode", "stacked")
self._preview_lock = threading.Lock() self._preview_lock = threading.Lock()
self._preview_busy = False self._preview_busy = False
@ -135,13 +412,16 @@ class RawSegformerService:
self.channels = int(cfg.get("channels", 4)) # 4 ou 5 self.channels = int(cfg.get("channels", 4)) # 4 ou 5
self.use_ndvi = bool(cfg.get("use_ndvi", False)) self.use_ndvi = bool(cfg.get("use_ndvi", False))
if self.channels not in (4, 5): if self.channels not in (3, 4, 5):
raise ValueError(f"[svc] 'channels' inválido no config: {self.channels} (esperado 4 ou 5)") raise ValueError(f"[svc] 'channels' inválido no config: {self.channels} (esperado 3, 4 ou 5)")
if self.channels == 5 and not self.use_ndvi: if self.channels == 5 and not self.use_ndvi:
raise ValueError("[svc] channels=5 sem NDVI não faz sentido no design atual. " raise ValueError("[svc] channels=5 sem NDVI não faz sentido no design atual. "
"Use channels=4 sem NDVI ou channels=5 com NDVI.") "Use channels=4 sem NDVI ou channels=5 com NDVI.")
if self.channels == 3 and self.use_ndvi:
raise ValueError("[svc] channels=3 com NDVI não faz sentido. Use RGB puro sem NDVI.")
# labelmap # labelmap
if labelmap_path is None: if labelmap_path is None:
dataset_path = os.path.join(self.MODELO, "dataset") dataset_path = os.path.join(self.MODELO, "dataset")
@ -164,8 +444,8 @@ class RawSegformerService:
# checkpoint # checkpoint
if ckpt_path is None: if ckpt_path is None:
# igual ao _9_test_segformer_b3_raw.py # igual ao _9_test_segformer_b3_raw.py
# save_path = MODELO/backup/modelo/model_name/rawx # save_path = MODELO/backup/modelo/model_name/fusion_rawx
save_path = os.path.join(self.MODELO, "backup", self.modelo_folder, self.MODEL_NAME, f"raw{self.channels}") save_path = os.path.join(self.MODELO, "backup", self.modelo_folder, self.MODEL_NAME, f"{self.fusion_mode}_raw{self.channels}")
ckpt_path = os.path.join(save_path, "best_miou.pt") ckpt_path = os.path.join(save_path, "best_miou.pt")
if not os.path.isfile(ckpt_path): if not os.path.isfile(ckpt_path):
@ -186,6 +466,7 @@ class RawSegformerService:
self._buf_hw = (H, W) self._buf_hw = (H, W)
# Buffers principais # Buffers principais
self._raw_buf3 = np.empty((3, H, W), dtype=np.float32)
self._raw_buf4 = np.empty((4, H, W), dtype=np.float32) self._raw_buf4 = np.empty((4, H, W), dtype=np.float32)
self._raw_buf5 = np.empty((5, H, W), dtype=np.float32) self._raw_buf5 = np.empty((5, H, W), dtype=np.float32)
@ -216,15 +497,28 @@ class RawSegformerService:
# Carregamento do modelo # Carregamento do modelo
# ------------------------------------------------- # -------------------------------------------------
def _load_model(self, backbone: str) -> torch.nn.Module: def _load_model(self, backbone: str) -> torch.nn.Module:
model = build_raw_segformer_model( if self.fusion_mode == "dual_branch":
num_classes=self.num_classes, model = build_dual_branch_segformer_model(
channels=self.channels, num_classes=self.num_classes,
backbone=backbone, channels=self.channels,
device=self.device, backbone=backbone,
ckpt_path=self.ckpt_path, device=self.device,
strict=True, ckpt_path=self.ckpt_path,
) strict=True,
print(f"[svc] modelo carregado: {self.ckpt_path}") )
elif self.fusion_mode == "stacked":
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,
)
else:
raise ValueError(f"[svc] fusion_mode inválido: {self.fusion_mode}")
print(f"[svc] modelo carregado: {self.ckpt_path} | fusion_mode={self.fusion_mode}")
return model return model
# ------------------------------------------------- # -------------------------------------------------
@ -248,7 +542,7 @@ class RawSegformerService:
Se você guardar a referência e chamar de novo, o conteúdo será sobrescrito. Se você guardar a referência e chamar de novo, o conteúdo será sobrescrito.
(No loop de tempo real isso é perfeito.) (No loop de tempo real isso é perfeito.)
""" """
if self.channels not in (4, 5): if self.channels not in (3, 4, 5):
raise RuntimeError(f"channels inválido em build_raw_input: {self.channels}") raise RuntimeError(f"channels inválido em build_raw_input: {self.channels}")
# Garantir float32 sem copiar se já estiver # Garantir float32 sem copiar se já estiver
@ -260,6 +554,13 @@ class RawSegformerService:
H, W = r32.shape H, W = r32.shape
self._ensure_buffers(H, W) self._ensure_buffers(H, W)
if self.channels == 3:
raw = self._raw_buf3
raw[0, :, :] = r32
raw[1, :, :] = g32
raw[2, :, :] = b32
return raw
# Seleciona buffer principal # Seleciona buffer principal
raw = self._raw_buf5 if self.channels == 5 else self._raw_buf4 raw = self._raw_buf5 if self.channels == 5 else self._raw_buf4
@ -550,12 +851,14 @@ class RawSegDataset(Dataset):
self.ndvi_eps = ndvi_eps self.ndvi_eps = ndvi_eps
self.channels = int(channels) self.channels = int(channels)
if self.channels not in (4, 5): if self.channels not in (3, 4, 5):
raise ValueError(f"[dataset] 'channels' inválido: {self.channels} (esperado 4 ou 5)") raise ValueError(f"[dataset] 'channels' inválido: {self.channels} (esperado 3, 4 ou 5)")
if self.channels == 5 and not self.use_ndvi: if self.channels == 5 and not self.use_ndvi:
raise ValueError("[dataset] channels=5 sem NDVI não é suportado. " raise ValueError("[svc] channels=5 sem NDVI não faz sentido no design atual.")
"Use channels=4 sem NDVI ou channels=5 com NDVI.")
if self.channels == 3 and self.use_ndvi:
raise ValueError("[svc] channels=3 com NDVI não faz sentido. Use RGB puro sem NDVI.")
self.items: List[Tuple[str, str]] = [] # (raw_path, mask_path) self.items: List[Tuple[str, str]] = [] # (raw_path, mask_path)
@ -742,16 +1045,21 @@ class RawSegDataset(Dataset):
# Montagem de canais finais # Montagem de canais finais
if self.use_ndvi: if self.use_ndvi:
if self.channels == 4: if self.channels == 4:
# substitui IR por NDVI01: [R,G,B,NDVI01] # [R,G,B,NDVI01]
raw = self._add_ndvi_replace(raw) raw = self._add_ndvi_replace(raw)
elif self.channels == 5: elif self.channels == 5:
# adiciona NDVI: [R,G,B,IR,NDVI01] # [R,G,B,IR,NDVI01]
raw = self._add_ndvi_append(raw) raw = self._add_ndvi_append(raw)
else: else:
raise RuntimeError(f"[dataset] channels inválido em __getitem__: {self.channels}") raise RuntimeError(f"[dataset] channels inválido com NDVI: {self.channels}")
else: else:
# modo clássico: [R,G,IR,B] if self.channels == 3:
if self.channels != 4: # [R,G,B]
raw = raw[..., :3]
elif self.channels == 4:
# [R,G,B,IR]
pass
else:
raise RuntimeError("[dataset] channels=5 sem NDVI não suportado em __getitem__.") raise RuntimeError("[dataset] channels=5 sem NDVI não suportado em __getitem__.")
mask = self._imread_mask_png(mask_path) # (H,W) uint8 mask = self._imread_mask_png(mask_path) # (H,W) uint8
@ -841,22 +1149,40 @@ def normalize_raw(imgs: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
return (imgs - mean) / std 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: def make_bgr_preview_from_raw(
raw_np: np.ndarray,
rgirb: bool,
preview_fast: bool,
preview_scale: int = 2,
apply_ir_comp: bool = True,
ir_k_r: float = 0.40,
ir_k_g: float = 0.10,
ir_k_b: float = 0.50,
) -> np.ndarray:
""" """
raw_np: (C,H,W) float32 em 0..1, C=4 ou 5. Gera preview BGR (OpenCV) a partir do tensor raw_np (C,H,W) float32 em 0..1.
Layout esperado: Casos esperados:
- channels=4, use_ndvi=False -> [R,G,B,IR] - rgirb=False:
- channels=4, use_ndvi=True -> [R,G,B,NDVI] raw_np = [R, G, B] ou [R, G, B, NDVI]
- channels=5, use_ndvi=True -> [R,G,B,IR,NDVI] - rgirb=True:
raw_np = [R, G, IR, B] ou [R, G, IR, B, NDVI]
Usamos sempre R,G,B para o preview. Regras:
- Preview usa sempre R,G,B
- Compensação de IR é aplicada quando rgirb=True
- preview_fast prioriza velocidade
- modo full prioriza legibilidade para anotação/mapeamento
""" """
assert raw_np.ndim == 3, "raw_np deve ser (C,H,W)" assert raw_np.ndim == 3, "raw_np deve ser (C,H,W)"
C = raw_np.shape[0] C = raw_np.shape[0]
if C not in (4, 5): if C not in (3, 4, 5):
raise RuntimeError(f"Esperado C=4 ou 5, veio {C} em make_rgb_preview_from_raw") raise RuntimeError(f"Esperado C=3, 4 ou 5, veio {C} em make_bgr_preview_from_raw")
# Garante float32 sem cópia desnecessária
raw_np = raw_np.astype(np.float32, copy=False)
# Seleção dos canais visíveis
r = raw_np[0] r = raw_np[0]
g = raw_np[1] g = raw_np[1]
b = raw_np[3] if rgirb else raw_np[2] b = raw_np[3] if rgirb else raw_np[2]
@ -865,52 +1191,65 @@ def make_bgr_preview_from_raw(raw_np: np.ndarray, rgirb: bool, preview_fast: boo
# Modo PREVIEW RÁPIDO # Modo PREVIEW RÁPIDO
# ------------------------- # -------------------------
if preview_fast: if preview_fast:
# Downscale simples (barato) # Downscale primeiro para baratear tudo depois
if preview_scale > 1: if preview_scale > 1:
r = r[::preview_scale, ::preview_scale] r = r[::preview_scale, ::preview_scale]
g = g[::preview_scale, ::preview_scale] g = g[::preview_scale, ::preview_scale]
b = b[::preview_scale, ::preview_scale] b = b[::preview_scale, ::preview_scale]
# Compensação leve de IR apenas se houver layout RGIRB
if rgirb and apply_ir_comp and C >= 4:
ir = raw_np[2]
if preview_scale > 1:
ir = ir[::preview_scale, ::preview_scale]
r = np.clip(r - ir_k_r * ir, 0.0, 1.0)
g = np.clip(g - ir_k_g * ir, 0.0, 1.0)
b = np.clip(b - ir_k_b * ir, 0.0, 1.0)
bgr = np.stack([b, g, r], axis=0) bgr = np.stack([b, g, r], axis=0)
# Gamma simples só pra não ficar "lavado" # Gamma simples para não ficar "lavado"
bgr = np.power(bgr, 1/1.8) bgr = np.power(np.clip(bgr, 0.0, 1.0), 1 / 1.8)
# Converte direto pra uint8
bgr8 = (bgr * 255.0).clip(0, 255).astype(np.uint8) bgr8 = (bgr * 255.0).clip(0, 255).astype(np.uint8)
bgr_hwc = np.transpose(bgr8, (1, 2, 0)) # RGB return np.transpose(bgr8, (1, 2, 0)).copy()
bgr = bgr_hwc.copy()
return bgr # -------------------------
# Modo FULL
# -------------------------
# Compensação de IR apenas se houver layout RGIRB
if rgirb and apply_ir_comp and C >= 4:
ir = raw_np[2]
r = np.clip(r - ir_k_r * ir, 0.0, 1.0)
g = np.clip(g - ir_k_g * ir, 0.0, 1.0)
b = np.clip(b - ir_k_b * ir, 0.0, 1.0)
bgr = np.stack([b, g, r], axis=0).astype(np.float32) # (3,H,W) def stretch_channel(x: np.ndarray, p_low: float = 1.0, p_high: float = 99.0) -> np.ndarray:
"""
Stretch robusto por percentil, melhor que AWB simples para cena agrícola.
"""
lo = np.percentile(x, p_low)
hi = np.percentile(x, p_high)
# Auto white balance simples (referenciado no G) if not np.isfinite(lo) or not np.isfinite(hi) or hi <= lo:
means = bgr.mean(axis=(1, 2)) # (3,) return np.clip(x, 0.0, 1.0)
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 x = (x - lo) / (hi - lo)
alpha = 1.0 # teste 0.4..0.8 return np.clip(x, 0.0, 1.0)
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 r = stretch_channel(r)
bgr = np.power(bgr, 1/1.9, out=bgr) g = stretch_channel(g)
b = stretch_channel(b)
# Para OpenCV bgr = np.stack([b, g, r], axis=0).astype(np.float32)
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 # Gamma suave
bgr = np.power(np.clip(bgr, 0.0, 1.0), 1 / 2.0)
bgr8 = (bgr * 255.0).clip(0, 255).astype(np.uint8)
return np.transpose(bgr8, (1, 2, 0)).copy()
# ---------------------------- # ----------------------------
# Patch: SegFormer para N canais # Patch: SegFormer para N canais