| """
|
| 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,
|
| "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]
|
|
|
|
|
|
|
|
|
| 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),
|
| ])
|
|
|
|
|
|
|
|
|
|
|
| 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)
|
|
|
|
|
| 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)
|
|
|
|
|
|
|
|
|
| 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()
|
|
|
|
|
|
|
|
|
|
|
| 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()}")))
|
|
|
| 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:
|
|
|
| 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]
|
|
|
|
|
| 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() |