Download test.py from UserPollo/moire-pattern-detector: direct link, hf CLI and curl.
- Browser
- Download file 16.5 kB
-
https://huggingface.co/UserPollo/moire-pattern-detector/resolve/refs%2Fpr%2F1/test.py
- Command line
-
hf download hf://UserPollo/moire-pattern-detector@refs/pr/1/test.py
-
curl -L -o test.py https://huggingface.co/UserPollo/moire-pattern-detector/resolve/refs%2Fpr%2F1/test.py
16.5 kB
| """ | |
| Evalúa el detector de moiré/pantalla (arquitectura dual-branch: local + | |
| global, backbone parcialmente descongelado) sobre: | |
| 1. El split "test" del dataset original (accuracy, matriz de confusión) | |
| 2. Un conjunto de fotos propias (sin etiquetas, solo predicción + confianza) | |
| Tiene que reproducir EXACTAMENTE la extracción de features del entrenamiento: | |
| CLS + promedio de patch tokens (sin register tokens) de la rama local | |
| (crop nativo) concatenado con lo mismo de la rama global (resize + crop). | |
| Uso: | |
| python test.py --images ruta/a/mis_fotos | |
| python test.py --test-dataset | |
| python test.py --images ruta/a/mis_fotos --test-dataset | |
| python test.py --test-dataset --tta # multi-crop, más lento pero más preciso | |
| Requiere que ya hayas corrido el script de entrenamiento, que deja en el | |
| directorio de trabajo: best_screen_detector_mlp.pt, classes.json, y | |
| (si unfreeze_last_n_blocks > 0) best_screen_detector_backbone.pt. | |
| """ | |
| import argparse | |
| import glob | |
| import json | |
| import os | |
| import random | |
| import torch | |
| import torch.nn as nn | |
| from PIL import Image | |
| from torchvision import transforms | |
| from torchvision.transforms import functional as TF | |
| from torch.utils.data import DataLoader, Dataset | |
| 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, | |
| "batch_size": 32, | |
| "tta_crops": 5, | |
| "ckpt_path": "best_screen_detector_mlp.pt", | |
| "backbone_ckpt_path": "best_screen_detector_backbone.pt", | |
| "classes_path": "classes.json", | |
| "dataset_slug": "soumikrakshit/uhdm-dataset", | |
| } | |
| 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 (deterministas, | |
| # sin augmentación): rama local = solo center crop sobre resolución nativa, | |
| # 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) | |
| def extract_dual_features(backbone, local_images, global_images): | |
| """Misma lógica que en el entrenamiento pero solo-inferencia: concatena | |
| local+global en el batch, un único forward, CLS + promedio de patch | |
| tokens (sin register tokens) de cada rama.""" | |
| 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, esos bloques | |
| # tienen pesos fine-tuneados guardados aparte -- hay que cargarlos, si | |
| # no, estaríamos evaluando con el backbone pre-entrenado original y | |
| # los resultados no coincidirían 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 " | |
| f"-- los resultados pueden no coincidir con el val_acc del entrenamiento.") | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| # Dataset de evaluación simple (una vista local + una vista global por imagen) | |
| # --------------------------------------------------------------------------- | |
| def _index_samples(root_dir): | |
| samples = [] | |
| for dirpath, _, filenames in os.walk(root_dir): | |
| for fname in filenames: | |
| lower = fname.lower() | |
| if not lower.endswith((".jpg", ".jpeg", ".png")): | |
| continue | |
| if "_gt" in lower: | |
| label = 0 | |
| elif "_moire" in lower: | |
| label = 1 | |
| else: | |
| continue | |
| samples.append((os.path.join(dirpath, fname), label)) | |
| if not samples: | |
| raise RuntimeError(f"No se encontraron imágenes '_gt'/'_moire' en {root_dir}") | |
| return samples | |
| class MoireDualDataset(Dataset): | |
| """Igual que en el entrenamiento: cada muestra devuelve (local, global, | |
| label), con transforms deterministas (sin augmentación) para evaluación.""" | |
| def __init__(self, root_dir): | |
| self.samples = _index_samples(root_dir) | |
| def __len__(self): | |
| return len(self.samples) | |
| def __getitem__(self, idx): | |
| path, label = self.samples[idx] | |
| img = Image.open(path).convert("RGB") | |
| local_img = local_eval_transform(img) | |
| global_img = global_eval_transform(img) | |
| return local_img, global_img, label | |
| # --------------------------------------------------------------------------- | |
| # Multi-crop TTA (opcional, --tta): mismo criterio que evaluate_tta() en el | |
| # script de entrenamiento -- centro + 4 esquinas como crops locales, más la | |
| # vista global, promediando las probabilidades. | |
| # --------------------------------------------------------------------------- | |
| 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) # (n_crops, C, H, W) | |
| 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) | |
| def predict_image_tta(backbone, head, img: Image.Image): | |
| """Predicción multi-crop para UNA imagen ya abierta (PIL). Devuelve | |
| (pred_idx, confidence, probs_tensor).""" | |
| 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() | |
| # --------------------------------------------------------------------------- | |
| # 1. Evaluación sobre el test set original | |
| # --------------------------------------------------------------------------- | |
| def evaluate_test_dataset(backbone, head, classes: list, use_tta: bool): | |
| import kagglehub | |
| dataset_root = kagglehub.dataset_download(CFG["dataset_slug"]) # cacheado | |
| test_dir = os.path.join(dataset_root, "test", "test") | |
| if not os.path.isdir(test_dir): | |
| print(f"Aviso: no encuentro {test_dir}.") | |
| return | |
| num_classes = len(classes) | |
| confusion = torch.zeros(num_classes, num_classes, dtype=torch.long) | |
| correct, total = 0, 0 | |
| if use_tta: | |
| samples = _index_samples(test_dir) | |
| print(f"\nEvaluando (TTA, {CFG['tta_crops']} crops/imagen) sobre " | |
| f"{len(samples)} imágenes del test set...") | |
| for path, label in samples: | |
| img = Image.open(path).convert("RGB") | |
| pred, _, _ = predict_image_tta(backbone, head, img) | |
| confusion[label, pred] += 1 | |
| correct += int(pred == label) | |
| total += 1 | |
| else: | |
| test_ds = MoireDualDataset(test_dir) | |
| loader = DataLoader(test_ds, batch_size=CFG["batch_size"], shuffle=False, | |
| num_workers=min(4, os.cpu_count() or 1), | |
| pin_memory=device.type == "cuda") | |
| print(f"\nEvaluando sobre {len(test_ds)} imágenes del test set...") | |
| for local_imgs, global_imgs, labels in loader: | |
| local_imgs, global_imgs = local_imgs.to(device), global_imgs.to(device) | |
| feats = extract_dual_features(backbone, local_imgs, global_imgs) | |
| logits = head(feats) | |
| pred = logits.argmax(dim=1).cpu() | |
| for t, p in zip(labels, pred): | |
| confusion[t, p] += 1 | |
| correct += (pred == labels).sum().item() | |
| total += labels.size(0) | |
| acc = 100 * correct / total | |
| print(f"\nAccuracy en test set: {acc:.2f}% ({correct}/{total})") | |
| print("\nMatriz de confusión (filas = real, columnas = predicho):") | |
| header = "".join(f"{c[:10]:>12}" for c in classes) | |
| print(" " * 12 + header) | |
| for i, row in enumerate(confusion): | |
| row_str = "".join(f"{v.item():>12}" for v in row) | |
| print(f"{classes[i][:10]:>12}{row_str}") | |
| print("\nPor clase:") | |
| for i, cls in enumerate(classes): | |
| tp = confusion[i, i].item() | |
| fn = confusion[i, :].sum().item() - tp | |
| fp = confusion[:, i].sum().item() - tp | |
| precision = tp / (tp + fp) if (tp + fp) else 0.0 | |
| recall = tp / (tp + fn) if (tp + fn) else 0.0 | |
| print(f" {cls}: precision={precision:.3f} recall={recall:.3f} (n={confusion[i, :].sum().item()})") | |
| # --------------------------------------------------------------------------- | |
| # 2. Predicción sobre fotos propias (sin etiquetas) | |
| # --------------------------------------------------------------------------- | |
| IMG_EXTENSIONS = (".jpg", ".jpeg", ".png", ".bmp", ".webp") | |
| 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()}"))) | |
| return sorted(paths) | |
| def predict_own_images(backbone, head, classes: list, images_path: str, use_tta: bool): | |
| 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"\nPrediciendo{tag} sobre {len(paths)} imagen(es) propias...") | |
| 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 | |
| if use_tta: | |
| pred_idx, conf, probs = predict_image_tta(backbone, head, img) | |
| else: | |
| local_img = local_eval_transform(img).unsqueeze(0).to(device) | |
| global_img = global_eval_transform(img).unsqueeze(0).to(device) | |
| feats = extract_dual_features(backbone, local_img, global_img) | |
| logits = head(feats) | |
| probs = torch.softmax(logits, dim=1)[0].cpu() | |
| conf, pred_idx = probs.max(dim=0) | |
| conf, pred_idx = conf.item(), pred_idx.item() | |
| pred_class = classes[pred_idx] | |
| print(f" {os.path.basename(path):<40} -> {pred_class:<15} " | |
| f"(confianza {conf*100:.1f}%) " | |
| f"[{', '.join(f'{c}={p*100:.1f}%' for c, p in zip(classes, probs.tolist()))}]") | |
| # --------------------------------------------------------------------------- | |
| # Main | |
| # --------------------------------------------------------------------------- | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Evalúa el detector de moiré/pantalla (dual-branch).") | |
| parser.add_argument("--images", type=str, default=None, | |
| help="Carpeta (o archivo) con tus propias fotos a evaluar.") | |
| parser.add_argument("--test-dataset", action="store_true", | |
| help="Evalúa sobre el split de test del dataset original, con métricas.") | |
| parser.add_argument("--tta", action="store_true", | |
| help="Usa multi-crop test-time augmentation (más lento, más preciso).") | |
| args = parser.parse_args() | |
| if not args.images and not args.test_dataset: | |
| parser.error("Especifica --images, --test-dataset, o ambos.") | |
| classes = load_classes() | |
| backbone, head = load_models(num_classes=len(classes)) | |
| if args.images: | |
| predict_own_images(backbone, head, classes, args.images, use_tta=args.tta) | |
| if args.test_dataset: | |
| evaluate_test_dataset(backbone, head, classes, use_tta=args.tta) | |
| if __name__ == "__main__": | |
| main() |