""" Inferencia del detector de moiré/pantalla (dual-branch: local + global, backbone parcialmente descongelado). Le pasás una carpeta (o un archivo) con imágenes y te devuelve, para cada una, la clase predicha y el % de confianza. Guarda las imágenes resultantes en una carpeta con el texto dibujado encima y opcionalmente un CSV. Uso: python inference.py --images ruta/a/mis_fotos python inference.py --images ruta/a/una_foto.jpg python inference.py --images ruta/a/mis_fotos --csv resultados.csv python inference.py --images ruta/a/mis_fotos --outdir mis_resultados python inference.py --images ruta/a/mis_fotos --tta # más lento, más preciso """ import argparse import csv import glob import json import os import random import torch import torch.nn as nn from PIL import Image, ImageDraw, ImageFont from torchvision import transforms from torchvision.transforms import functional as TF from transformers import AutoModel CFG = { "backbone_name": "facebook/dinov2-with-registers-base", "input_size": 224, "global_resize": 256, "unfreeze_last_n_blocks": 2, # debe coincidir con lo usado en el entrenamiento "hidden_size": 256, "dropout": 0.3, "ckpt_path": "best_screen_detector_mlp.pt", "backbone_ckpt_path": "best_screen_detector_backbone.pt", "classes_path": "classes.json", "batch_size": 32, "tta_crops": 5, } IMG_EXTENSIONS = (".jpg", ".jpeg", ".png", ".bmp", ".webp") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] # Mismas transforms de evaluación que en el entrenamiento: rama local = solo # center crop sobre resolución nativa (sin destruir el moiré con un resize # previo), rama global = resize completo + center crop. local_eval_transform = transforms.Compose([ transforms.CenterCrop(CFG["input_size"]), transforms.ToTensor(), transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ]) global_eval_transform = transforms.Compose([ transforms.Resize((CFG["global_resize"], CFG["global_resize"])), transforms.CenterCrop(CFG["input_size"]), transforms.ToTensor(), transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ]) tta_base_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), ]) # --------------------------------------------------------------------------- # Modelo (misma arquitectura que en el entrenamiento) # --------------------------------------------------------------------------- class ScreenDetectorMLP(nn.Module): """input_size = 4 * hidden dim del backbone: (CLS + patch-mean) de la rama local concatenado con (CLS + patch-mean) de la rama global.""" def __init__(self, input_size: int = 3072, hidden_size: int = 256, num_classes: int = 2, dropout: float = 0.3): super().__init__() self.mlp = nn.Sequential( nn.Linear(input_size, hidden_size), nn.GELU(), nn.BatchNorm1d(hidden_size), nn.Dropout(dropout), nn.Linear(hidden_size, hidden_size // 2), nn.GELU(), nn.Dropout(dropout), nn.Linear(hidden_size // 2, num_classes), ) def forward(self, x): return self.mlp(x) def get_num_register_tokens(backbone) -> int: return getattr(backbone.config, "num_register_tokens", 0) @torch.no_grad() def extract_dual_features(backbone, local_images, global_images): """Misma lógica que en el entrenamiento: concatena local+global en el batch, un único forward, CLS + promedio de patch tokens (sin register tokens) de cada rama, concatenados.""" n_reg = get_num_register_tokens(backbone) batch = torch.cat([local_images, global_images], dim=0) if device.type == "cuda": with torch.autocast(device_type="cuda", dtype=torch.float16): out = backbone(pixel_values=batch) else: out = backbone(pixel_values=batch) hidden = out.last_hidden_state.float() cls_tok = hidden[:, 0, :] patch_mean = hidden[:, 1 + n_reg:, :].mean(dim=1) feat = torch.cat([cls_tok, patch_mean], dim=-1) B = local_images.size(0) local_feat, global_feat = feat[:B], feat[B:] return torch.cat([local_feat, global_feat], dim=-1) # (B, 4*hidden) def load_classes() -> list: if not os.path.exists(CFG["classes_path"]): raise FileNotFoundError( f"No encuentro {CFG['classes_path']}. Corre primero el script de entrenamiento." ) with open(CFG["classes_path"]) as f: return json.load(f) def load_models(num_classes: int): print(f"Cargando backbone {CFG['backbone_name']}...") backbone = AutoModel.from_pretrained(CFG["backbone_name"]).to(device) backbone.eval() for p in backbone.parameters(): p.requires_grad_(False) # Si el entrenamiento descongeló los últimos N bloques, hay que cargar # esos pesos fine-tuneados; si no, se evalúa con el backbone original # y los resultados no coinciden con el checkpoint de la cabeza. n_unfreeze = CFG["unfreeze_last_n_blocks"] if n_unfreeze > 0: if os.path.exists(CFG["backbone_ckpt_path"]): total_layers = len(backbone.encoder.layer) unfrozen_state = torch.load(CFG["backbone_ckpt_path"], map_location=device) for i, layer in enumerate(backbone.encoder.layer[total_layers - n_unfreeze:]): layer.load_state_dict(unfrozen_state[f"layer.{total_layers - n_unfreeze + i}"]) print(f"Pesos fine-tuneados de los últimos {n_unfreeze} bloques cargados " f"desde {CFG['backbone_ckpt_path']}") else: print(f"AVISO: unfreeze_last_n_blocks={n_unfreeze} pero no existe " f"{CFG['backbone_ckpt_path']}. Evaluando con el backbone SIN fine-tunear.") feat_dim = backbone.config.hidden_size * 4 head = ScreenDetectorMLP(input_size=feat_dim, hidden_size=CFG["hidden_size"], num_classes=num_classes, dropout=CFG["dropout"]).to(device) head.load_state_dict(torch.load(CFG["ckpt_path"], map_location=device)) head.eval() print(f"Pesos del MLP cargados desde {CFG['ckpt_path']}") return backbone, head @torch.no_grad() def predict_batch(backbone, head, local_images: torch.Tensor, global_images: torch.Tensor): local_images = local_images.to(device, non_blocking=True) global_images = global_images.to(device, non_blocking=True) feats = extract_dual_features(backbone, local_images, global_images) logits = head(feats) probs = torch.softmax(logits, dim=1) conf, pred = probs.max(dim=1) return pred.cpu(), conf.cpu(), probs.cpu() # --------------------------------------------------------------------------- # Multi-crop TTA (opcional, --tta) # --------------------------------------------------------------------------- def make_local_crops(img: Image.Image, n_crops: int, crop_size: int) -> torch.Tensor: w, h = img.size cs = crop_size cx, cy = max((w - cs) // 2, 0), max((h - cs) // 2, 0) positions = [(cx, cy), (0, 0), (max(w - cs, 0), 0), (0, max(h - cs, 0)), (max(w - cs, 0), max(h - cs, 0))] while len(positions) < n_crops: positions.append((random.randint(0, max(w - cs, 0)), random.randint(0, max(h - cs, 0)))) positions = positions[:n_crops] crops = [] for x, y in positions: crop = img.crop((x, y, x + cs, y + cs)) if crop.size != (cs, cs): crop = crop.resize((cs, cs)) crops.append(tta_base_transform(crop)) return torch.stack(crops) def make_global_view(img: Image.Image, global_resize: int, crop_size: int) -> torch.Tensor: g = TF.resize(img, [global_resize, global_resize]) g = TF.center_crop(g, [crop_size, crop_size]) return tta_base_transform(g) @torch.no_grad() def predict_image_tta(backbone, head, img: Image.Image): local_crops = make_local_crops(img, CFG["tta_crops"], CFG["input_size"]).to(device) global_view = make_global_view(img, CFG["global_resize"], CFG["input_size"]) global_crops = global_view.unsqueeze(0).expand(CFG["tta_crops"], -1, -1, -1).contiguous().to(device) feats = extract_dual_features(backbone, local_crops, global_crops) logits = head(feats) probs = torch.softmax(logits, dim=-1).mean(dim=0) conf, pred = probs.max(dim=0) return pred.item(), conf.item(), probs.cpu() def collect_image_paths(path: str) -> list: if os.path.isfile(path): return [path] paths = [] for ext in IMG_EXTENSIONS: paths.extend(glob.glob(os.path.join(path, f"*{ext}"))) paths.extend(glob.glob(os.path.join(path, f"*{ext.upper()}"))) # también busca en subcarpetas, por si la organización no es plana for ext in IMG_EXTENSIONS: paths.extend(glob.glob(os.path.join(path, "**", f"*{ext}"), recursive=True)) paths.extend(glob.glob(os.path.join(path, "**", f"*{ext.upper()}"), recursive=True)) return sorted(set(paths)) def load_image_batch(paths: list): """Carga y transforma un batch de imágenes (ambas ramas); descarta las que fallen al abrir.""" local_tensors, global_tensors, valid_paths = [], [], [] for path in paths: try: img = Image.open(path).convert("RGB") except Exception as e: print(f" {os.path.basename(path)}: no se pudo abrir ({e})") continue local_tensors.append(local_eval_transform(img)) global_tensors.append(global_eval_transform(img)) valid_paths.append(path) if not local_tensors: return None, None, [] return torch.stack(local_tensors), torch.stack(global_tensors), valid_paths def run_inference(backbone, head, classes: list, images_path: str, csv_path: str = None, only_class: str = None, out_dir: str = "results", use_tta: bool = False): paths = collect_image_paths(images_path) if not paths: print(f"No encontré imágenes en {images_path}") return tag = " (TTA)" if use_tta else "" print(f"\nProcesando{tag} {len(paths)} imagen(es)...\n") results = [] os.makedirs(out_dir, exist_ok=True) if use_tta: # TTA es por-imagen (5 forwards c/u), no se batchea entre imágenes for path in paths: try: img = Image.open(path).convert("RGB") except Exception as e: print(f" {os.path.basename(path)}: no se pudo abrir ({e})") continue pred_idx, conf, probs = predict_image_tta(backbone, head, img) results.append({ "file": path, "pred_class": classes[pred_idx], "confidence": conf * 100, **{cls: probs[j].item() * 100 for j, cls in enumerate(classes)}, }) else: batch_size = CFG["batch_size"] for i in range(0, len(paths), batch_size): chunk = paths[i:i + batch_size] local_batch, global_batch, valid_paths = load_image_batch(chunk) if local_batch is None: continue pred, conf, probs = predict_batch(backbone, head, local_batch, global_batch) for path, p, c, pr in zip(valid_paths, pred, conf, probs): results.append({ "file": path, "pred_class": classes[p.item()], "confidence": c.item() * 100, **{cls: pr[j].item() * 100 for j, cls in enumerate(classes)}, }) if only_class: results = [r for r in results if r["pred_class"] == only_class] # orden: menor confianza primero, para que lo más dudoso salte a la vista results.sort(key=lambda r: r["confidence"]) try: font = ImageFont.truetype("arial.ttf", 36) except IOError: font = ImageFont.load_default() print(f"\nGuardando imágenes anotadas en la carpeta '{out_dir}/'...") for r in results: detail = ", ".join(f"{cls}={r[cls]:.1f}%" for cls in classes) print(f" {os.path.basename(r['file']):<40} -> {r['pred_class']:<10} " f"(confianza {r['confidence']:.1f}%) [{detail}]") try: img = Image.open(r["file"]).convert("RGB") draw = ImageDraw.Draw(img) text = f"{r['pred_class']}: {r['confidence']:.1f}%" bbox = draw.textbbox((10, 10), text, font=font) draw.rectangle([bbox[0] - 5, bbox[1] - 5, bbox[2] + 5, bbox[3] + 5], fill="black") draw.text((10, 10), text, fill="white", font=font) save_path = os.path.join(out_dir, os.path.basename(r["file"])) img.save(save_path) except Exception as e: print(f"No se pudo procesar y guardar la imagen {r['file']}: {e}") print(f"\nTotal: {len(results)} imagen(es) predicha(s).") if classes: for cls in classes: n = sum(1 for r in results if r["pred_class"] == cls) print(f" {cls}: {n}") if csv_path: fieldnames = ["file", "pred_class", "confidence"] + classes with open(csv_path, "w", newline="") as f: writer = csv.DictWriter(f, fieldnames=fieldnames) writer.writeheader() for r in results: writer.writerow(r) print(f"\nResultados guardados en {csv_path}") def main(): parser = argparse.ArgumentParser(description="Predice moiré/pantalla sobre tus propias imágenes.") parser.add_argument("--images", type=str, required=True, help="Carpeta (o archivo) con las imágenes a evaluar.") parser.add_argument("--csv", type=str, default=None, help="Ruta opcional para guardar los resultados en CSV.") parser.add_argument("--only", type=str, default=None, help="Mostrar solo las imágenes predichas con esta clase (ej: moire).") parser.add_argument("--outdir", type=str, default="results", help="Carpeta donde se guardarán las imágenes con el resultado (por defecto 'results').") parser.add_argument("--tta", action="store_true", help="Usa multi-crop test-time augmentation (más lento, más preciso).") args = parser.parse_args() classes = load_classes() backbone, head = load_models(num_classes=len(classes)) run_inference(backbone, head, classes, args.images, csv_path=args.csv, only_class=args.only, out_dir=args.outdir, use_tta=args.tta) if __name__ == "__main__": main()