UserPollo commited on
Commit
022cb4d
·
verified ·
1 Parent(s): fc666d4

Necessary files

Browse files
best_screen_detector_backbone.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6a9e963dba8c9ece822c932fbdc3b501a50cdba94365b98625b6e0a48a6be9a0
3
+ size 56729299
best_screen_detector_mlp.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:73c6656bc7941049dabe3f274c7992d333a792b9ca292701a7ac143110536224
3
+ size 3288419
classes.json ADDED
@@ -0,0 +1 @@
 
 
1
+ ["gt", "moire"]
inference.py ADDED
@@ -0,0 +1,369 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Inferencia del detector de moiré/pantalla (dual-branch: local + global,
3
+ backbone parcialmente descongelado).
4
+
5
+ Le pasás una carpeta (o un archivo) con imágenes y te devuelve, para cada
6
+ una, la clase predicha y el % de confianza. Guarda las imágenes resultantes
7
+ en una carpeta con el texto dibujado encima y opcionalmente un CSV.
8
+
9
+ Uso:
10
+ python inference.py --images ruta/a/mis_fotos
11
+ python inference.py --images ruta/a/una_foto.jpg
12
+ python inference.py --images ruta/a/mis_fotos --csv resultados.csv
13
+ python inference.py --images ruta/a/mis_fotos --outdir mis_resultados
14
+ python inference.py --images ruta/a/mis_fotos --tta # más lento, más preciso
15
+ """
16
+
17
+ import argparse
18
+ import csv
19
+ import glob
20
+ import json
21
+ import os
22
+ import random
23
+
24
+ import torch
25
+ import torch.nn as nn
26
+ from PIL import Image, ImageDraw, ImageFont
27
+ from torchvision import transforms
28
+ from torchvision.transforms import functional as TF
29
+ from transformers import AutoModel
30
+
31
+ CFG = {
32
+ "backbone_name": "facebook/dinov2-with-registers-base",
33
+ "input_size": 224,
34
+ "global_resize": 256,
35
+ "unfreeze_last_n_blocks": 2, # debe coincidir con lo usado en el entrenamiento
36
+ "hidden_size": 256,
37
+ "dropout": 0.3,
38
+ "ckpt_path": "best_screen_detector_mlp.pt",
39
+ "backbone_ckpt_path": "best_screen_detector_backbone.pt",
40
+ "classes_path": "classes.json",
41
+ "batch_size": 32,
42
+ "tta_crops": 5,
43
+ }
44
+
45
+ IMG_EXTENSIONS = (".jpg", ".jpeg", ".png", ".bmp", ".webp")
46
+
47
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
48
+
49
+ IMAGENET_MEAN = [0.485, 0.456, 0.406]
50
+ IMAGENET_STD = [0.229, 0.224, 0.225]
51
+
52
+ # Mismas transforms de evaluación que en el entrenamiento: rama local = solo
53
+ # center crop sobre resolución nativa (sin destruir el moiré con un resize
54
+ # previo), rama global = resize completo + center crop.
55
+ local_eval_transform = transforms.Compose([
56
+ transforms.CenterCrop(CFG["input_size"]),
57
+ transforms.ToTensor(),
58
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
59
+ ])
60
+
61
+ global_eval_transform = transforms.Compose([
62
+ transforms.Resize((CFG["global_resize"], CFG["global_resize"])),
63
+ transforms.CenterCrop(CFG["input_size"]),
64
+ transforms.ToTensor(),
65
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
66
+ ])
67
+
68
+ tta_base_transform = transforms.Compose([
69
+ transforms.ToTensor(),
70
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
71
+ ])
72
+
73
+
74
+ # ---------------------------------------------------------------------------
75
+ # Modelo (misma arquitectura que en el entrenamiento)
76
+ # ---------------------------------------------------------------------------
77
+ class ScreenDetectorMLP(nn.Module):
78
+ """input_size = 4 * hidden dim del backbone: (CLS + patch-mean) de la
79
+ rama local concatenado con (CLS + patch-mean) de la rama global."""
80
+
81
+ def __init__(self, input_size: int = 3072, hidden_size: int = 256,
82
+ num_classes: int = 2, dropout: float = 0.3):
83
+ super().__init__()
84
+ self.mlp = nn.Sequential(
85
+ nn.Linear(input_size, hidden_size),
86
+ nn.GELU(),
87
+ nn.BatchNorm1d(hidden_size),
88
+ nn.Dropout(dropout),
89
+ nn.Linear(hidden_size, hidden_size // 2),
90
+ nn.GELU(),
91
+ nn.Dropout(dropout),
92
+ nn.Linear(hidden_size // 2, num_classes),
93
+ )
94
+
95
+ def forward(self, x):
96
+ return self.mlp(x)
97
+
98
+
99
+ def get_num_register_tokens(backbone) -> int:
100
+ return getattr(backbone.config, "num_register_tokens", 0)
101
+
102
+
103
+ @torch.no_grad()
104
+ def extract_dual_features(backbone, local_images, global_images):
105
+ """Misma lógica que en el entrenamiento: concatena local+global en el
106
+ batch, un único forward, CLS + promedio de patch tokens (sin register
107
+ tokens) de cada rama, concatenados."""
108
+ n_reg = get_num_register_tokens(backbone)
109
+ batch = torch.cat([local_images, global_images], dim=0)
110
+
111
+ if device.type == "cuda":
112
+ with torch.autocast(device_type="cuda", dtype=torch.float16):
113
+ out = backbone(pixel_values=batch)
114
+ else:
115
+ out = backbone(pixel_values=batch)
116
+ hidden = out.last_hidden_state.float()
117
+
118
+ cls_tok = hidden[:, 0, :]
119
+ patch_mean = hidden[:, 1 + n_reg:, :].mean(dim=1)
120
+ feat = torch.cat([cls_tok, patch_mean], dim=-1)
121
+
122
+ B = local_images.size(0)
123
+ local_feat, global_feat = feat[:B], feat[B:]
124
+ return torch.cat([local_feat, global_feat], dim=-1) # (B, 4*hidden)
125
+
126
+
127
+ def load_classes() -> list:
128
+ if not os.path.exists(CFG["classes_path"]):
129
+ raise FileNotFoundError(
130
+ f"No encuentro {CFG['classes_path']}. Corre primero el script de entrenamiento."
131
+ )
132
+ with open(CFG["classes_path"]) as f:
133
+ return json.load(f)
134
+
135
+
136
+ def load_models(num_classes: int):
137
+ print(f"Cargando backbone {CFG['backbone_name']}...")
138
+ backbone = AutoModel.from_pretrained(CFG["backbone_name"]).to(device)
139
+ backbone.eval()
140
+ for p in backbone.parameters():
141
+ p.requires_grad_(False)
142
+
143
+ # Si el entrenamiento descongeló los últimos N bloques, hay que cargar
144
+ # esos pesos fine-tuneados; si no, se evalúa con el backbone original
145
+ # y los resultados no coinciden con el checkpoint de la cabeza.
146
+ n_unfreeze = CFG["unfreeze_last_n_blocks"]
147
+ if n_unfreeze > 0:
148
+ if os.path.exists(CFG["backbone_ckpt_path"]):
149
+ total_layers = len(backbone.encoder.layer)
150
+ unfrozen_state = torch.load(CFG["backbone_ckpt_path"], map_location=device)
151
+ for i, layer in enumerate(backbone.encoder.layer[total_layers - n_unfreeze:]):
152
+ layer.load_state_dict(unfrozen_state[f"layer.{total_layers - n_unfreeze + i}"])
153
+ print(f"Pesos fine-tuneados de los últimos {n_unfreeze} bloques cargados "
154
+ f"desde {CFG['backbone_ckpt_path']}")
155
+ else:
156
+ print(f"AVISO: unfreeze_last_n_blocks={n_unfreeze} pero no existe "
157
+ f"{CFG['backbone_ckpt_path']}. Evaluando con el backbone SIN fine-tunear.")
158
+
159
+ feat_dim = backbone.config.hidden_size * 4
160
+ head = ScreenDetectorMLP(input_size=feat_dim, hidden_size=CFG["hidden_size"],
161
+ num_classes=num_classes, dropout=CFG["dropout"]).to(device)
162
+ head.load_state_dict(torch.load(CFG["ckpt_path"], map_location=device))
163
+ head.eval()
164
+ print(f"Pesos del MLP cargados desde {CFG['ckpt_path']}")
165
+ return backbone, head
166
+
167
+
168
+ @torch.no_grad()
169
+ def predict_batch(backbone, head, local_images: torch.Tensor, global_images: torch.Tensor):
170
+ local_images = local_images.to(device, non_blocking=True)
171
+ global_images = global_images.to(device, non_blocking=True)
172
+ feats = extract_dual_features(backbone, local_images, global_images)
173
+ logits = head(feats)
174
+ probs = torch.softmax(logits, dim=1)
175
+ conf, pred = probs.max(dim=1)
176
+ return pred.cpu(), conf.cpu(), probs.cpu()
177
+
178
+
179
+ # ---------------------------------------------------------------------------
180
+ # Multi-crop TTA (opcional, --tta)
181
+ # ---------------------------------------------------------------------------
182
+ def make_local_crops(img: Image.Image, n_crops: int, crop_size: int) -> torch.Tensor:
183
+ w, h = img.size
184
+ cs = crop_size
185
+ cx, cy = max((w - cs) // 2, 0), max((h - cs) // 2, 0)
186
+ positions = [(cx, cy),
187
+ (0, 0), (max(w - cs, 0), 0), (0, max(h - cs, 0)), (max(w - cs, 0), max(h - cs, 0))]
188
+ while len(positions) < n_crops:
189
+ positions.append((random.randint(0, max(w - cs, 0)), random.randint(0, max(h - cs, 0))))
190
+ positions = positions[:n_crops]
191
+
192
+ crops = []
193
+ for x, y in positions:
194
+ crop = img.crop((x, y, x + cs, y + cs))
195
+ if crop.size != (cs, cs):
196
+ crop = crop.resize((cs, cs))
197
+ crops.append(tta_base_transform(crop))
198
+ return torch.stack(crops)
199
+
200
+
201
+ def make_global_view(img: Image.Image, global_resize: int, crop_size: int) -> torch.Tensor:
202
+ g = TF.resize(img, [global_resize, global_resize])
203
+ g = TF.center_crop(g, [crop_size, crop_size])
204
+ return tta_base_transform(g)
205
+
206
+
207
+ @torch.no_grad()
208
+ def predict_image_tta(backbone, head, img: Image.Image):
209
+ local_crops = make_local_crops(img, CFG["tta_crops"], CFG["input_size"]).to(device)
210
+ global_view = make_global_view(img, CFG["global_resize"], CFG["input_size"])
211
+ global_crops = global_view.unsqueeze(0).expand(CFG["tta_crops"], -1, -1, -1).contiguous().to(device)
212
+
213
+ feats = extract_dual_features(backbone, local_crops, global_crops)
214
+ logits = head(feats)
215
+ probs = torch.softmax(logits, dim=-1).mean(dim=0)
216
+ conf, pred = probs.max(dim=0)
217
+ return pred.item(), conf.item(), probs.cpu()
218
+
219
+
220
+ def collect_image_paths(path: str) -> list:
221
+ if os.path.isfile(path):
222
+ return [path]
223
+ paths = []
224
+ for ext in IMG_EXTENSIONS:
225
+ paths.extend(glob.glob(os.path.join(path, f"*{ext}")))
226
+ paths.extend(glob.glob(os.path.join(path, f"*{ext.upper()}")))
227
+ # también busca en subcarpetas, por si la organización no es plana
228
+ for ext in IMG_EXTENSIONS:
229
+ paths.extend(glob.glob(os.path.join(path, "**", f"*{ext}"), recursive=True))
230
+ paths.extend(glob.glob(os.path.join(path, "**", f"*{ext.upper()}"), recursive=True))
231
+ return sorted(set(paths))
232
+
233
+
234
+ def load_image_batch(paths: list):
235
+ """Carga y transforma un batch de imágenes (ambas ramas); descarta las
236
+ que fallen al abrir."""
237
+ local_tensors, global_tensors, valid_paths = [], [], []
238
+ for path in paths:
239
+ try:
240
+ img = Image.open(path).convert("RGB")
241
+ except Exception as e:
242
+ print(f" {os.path.basename(path)}: no se pudo abrir ({e})")
243
+ continue
244
+ local_tensors.append(local_eval_transform(img))
245
+ global_tensors.append(global_eval_transform(img))
246
+ valid_paths.append(path)
247
+ if not local_tensors:
248
+ return None, None, []
249
+ return torch.stack(local_tensors), torch.stack(global_tensors), valid_paths
250
+
251
+
252
+ def run_inference(backbone, head, classes: list, images_path: str,
253
+ csv_path: str = None, only_class: str = None, out_dir: str = "results",
254
+ use_tta: bool = False):
255
+ paths = collect_image_paths(images_path)
256
+ if not paths:
257
+ print(f"No encontré imágenes en {images_path}")
258
+ return
259
+
260
+ tag = " (TTA)" if use_tta else ""
261
+ print(f"\nProcesando{tag} {len(paths)} imagen(es)...\n")
262
+ results = []
263
+
264
+ os.makedirs(out_dir, exist_ok=True)
265
+
266
+ if use_tta:
267
+ # TTA es por-imagen (5 forwards c/u), no se batchea entre imágenes
268
+ for path in paths:
269
+ try:
270
+ img = Image.open(path).convert("RGB")
271
+ except Exception as e:
272
+ print(f" {os.path.basename(path)}: no se pudo abrir ({e})")
273
+ continue
274
+ pred_idx, conf, probs = predict_image_tta(backbone, head, img)
275
+ results.append({
276
+ "file": path,
277
+ "pred_class": classes[pred_idx],
278
+ "confidence": conf * 100,
279
+ **{cls: probs[j].item() * 100 for j, cls in enumerate(classes)},
280
+ })
281
+ else:
282
+ batch_size = CFG["batch_size"]
283
+ for i in range(0, len(paths), batch_size):
284
+ chunk = paths[i:i + batch_size]
285
+ local_batch, global_batch, valid_paths = load_image_batch(chunk)
286
+ if local_batch is None:
287
+ continue
288
+
289
+ pred, conf, probs = predict_batch(backbone, head, local_batch, global_batch)
290
+
291
+ for path, p, c, pr in zip(valid_paths, pred, conf, probs):
292
+ results.append({
293
+ "file": path,
294
+ "pred_class": classes[p.item()],
295
+ "confidence": c.item() * 100,
296
+ **{cls: pr[j].item() * 100 for j, cls in enumerate(classes)},
297
+ })
298
+
299
+ if only_class:
300
+ results = [r for r in results if r["pred_class"] == only_class]
301
+
302
+ # orden: menor confianza primero, para que lo más dudoso salte a la vista
303
+ results.sort(key=lambda r: r["confidence"])
304
+
305
+ try:
306
+ font = ImageFont.truetype("arial.ttf", 36)
307
+ except IOError:
308
+ font = ImageFont.load_default()
309
+
310
+ print(f"\nGuardando imágenes anotadas en la carpeta '{out_dir}/'...")
311
+ for r in results:
312
+ detail = ", ".join(f"{cls}={r[cls]:.1f}%" for cls in classes)
313
+ print(f" {os.path.basename(r['file']):<40} -> {r['pred_class']:<10} "
314
+ f"(confianza {r['confidence']:.1f}%) [{detail}]")
315
+
316
+ try:
317
+ img = Image.open(r["file"]).convert("RGB")
318
+ draw = ImageDraw.Draw(img)
319
+
320
+ text = f"{r['pred_class']}: {r['confidence']:.1f}%"
321
+
322
+ bbox = draw.textbbox((10, 10), text, font=font)
323
+ draw.rectangle([bbox[0] - 5, bbox[1] - 5, bbox[2] + 5, bbox[3] + 5], fill="black")
324
+ draw.text((10, 10), text, fill="white", font=font)
325
+
326
+ save_path = os.path.join(out_dir, os.path.basename(r["file"]))
327
+ img.save(save_path)
328
+
329
+ except Exception as e:
330
+ print(f"No se pudo procesar y guardar la imagen {r['file']}: {e}")
331
+
332
+ print(f"\nTotal: {len(results)} imagen(es) predicha(s).")
333
+ if classes:
334
+ for cls in classes:
335
+ n = sum(1 for r in results if r["pred_class"] == cls)
336
+ print(f" {cls}: {n}")
337
+
338
+ if csv_path:
339
+ fieldnames = ["file", "pred_class", "confidence"] + classes
340
+ with open(csv_path, "w", newline="") as f:
341
+ writer = csv.DictWriter(f, fieldnames=fieldnames)
342
+ writer.writeheader()
343
+ for r in results:
344
+ writer.writerow(r)
345
+ print(f"\nResultados guardados en {csv_path}")
346
+
347
+
348
+ def main():
349
+ parser = argparse.ArgumentParser(description="Predice moiré/pantalla sobre tus propias imágenes.")
350
+ parser.add_argument("--images", type=str, required=True,
351
+ help="Carpeta (o archivo) con las imágenes a evaluar.")
352
+ parser.add_argument("--csv", type=str, default=None,
353
+ help="Ruta opcional para guardar los resultados en CSV.")
354
+ parser.add_argument("--only", type=str, default=None,
355
+ help="Mostrar solo las imágenes predichas con esta clase (ej: moire).")
356
+ parser.add_argument("--outdir", type=str, default="results",
357
+ help="Carpeta donde se guardarán las imágenes con el resultado (por defecto 'results').")
358
+ parser.add_argument("--tta", action="store_true",
359
+ help="Usa multi-crop test-time augmentation (más lento, más preciso).")
360
+ args = parser.parse_args()
361
+
362
+ classes = load_classes()
363
+ backbone, head = load_models(num_classes=len(classes))
364
+ run_inference(backbone, head, classes, args.images, csv_path=args.csv,
365
+ only_class=args.only, out_dir=args.outdir, use_tta=args.tta)
366
+
367
+
368
+ if __name__ == "__main__":
369
+ main()
test.py ADDED
@@ -0,0 +1,397 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Evalúa el detector de moiré/pantalla (arquitectura dual-branch: local +
3
+ global, backbone parcialmente descongelado) sobre:
4
+ 1. El split "test" del dataset original (accuracy, matriz de confusión)
5
+ 2. Un conjunto de fotos propias (sin etiquetas, solo predicción + confianza)
6
+
7
+ Tiene que reproducir EXACTAMENTE la extracción de features del entrenamiento:
8
+ CLS + promedio de patch tokens (sin register tokens) de la rama local
9
+ (crop nativo) concatenado con lo mismo de la rama global (resize + crop).
10
+
11
+ Uso:
12
+ python test.py --images ruta/a/mis_fotos
13
+ python test.py --test-dataset
14
+ python test.py --images ruta/a/mis_fotos --test-dataset
15
+ python test.py --test-dataset --tta # multi-crop, más lento pero más preciso
16
+
17
+ Requiere que ya hayas corrido el script de entrenamiento, que deja en el
18
+ directorio de trabajo: best_screen_detector_mlp.pt, classes.json, y
19
+ (si unfreeze_last_n_blocks > 0) best_screen_detector_backbone.pt.
20
+ """
21
+
22
+ import argparse
23
+ import glob
24
+ import json
25
+ import os
26
+ import random
27
+
28
+ import torch
29
+ import torch.nn as nn
30
+ from PIL import Image
31
+ from torchvision import transforms
32
+ from torchvision.transforms import functional as TF
33
+ from torch.utils.data import DataLoader, Dataset
34
+ from transformers import AutoModel
35
+
36
+ CFG = {
37
+ "backbone_name": "facebook/dinov2-with-registers-base",
38
+ "input_size": 224,
39
+ "global_resize": 256,
40
+ "unfreeze_last_n_blocks": 2, # debe coincidir con lo usado en el entrenamiento
41
+ "hidden_size": 256,
42
+ "dropout": 0.3,
43
+ "batch_size": 32,
44
+ "tta_crops": 5,
45
+ "ckpt_path": "best_screen_detector_mlp.pt",
46
+ "backbone_ckpt_path": "best_screen_detector_backbone.pt",
47
+ "classes_path": "classes.json",
48
+ "dataset_slug": "soumikrakshit/uhdm-dataset",
49
+ }
50
+
51
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
52
+
53
+ IMAGENET_MEAN = [0.485, 0.456, 0.406]
54
+ IMAGENET_STD = [0.229, 0.224, 0.225]
55
+
56
+ # Mismas transforms de evaluación que en el entrenamiento (deterministas,
57
+ # sin augmentación): rama local = solo center crop sobre resolución nativa,
58
+ # rama global = resize completo + center crop.
59
+ local_eval_transform = transforms.Compose([
60
+ transforms.CenterCrop(CFG["input_size"]),
61
+ transforms.ToTensor(),
62
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
63
+ ])
64
+
65
+ global_eval_transform = transforms.Compose([
66
+ transforms.Resize((CFG["global_resize"], CFG["global_resize"])),
67
+ transforms.CenterCrop(CFG["input_size"]),
68
+ transforms.ToTensor(),
69
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
70
+ ])
71
+
72
+ tta_base_transform = transforms.Compose([
73
+ transforms.ToTensor(),
74
+ transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
75
+ ])
76
+
77
+
78
+ # ---------------------------------------------------------------------------
79
+ # Modelo (misma arquitectura que en el entrenamiento)
80
+ # ---------------------------------------------------------------------------
81
+ class ScreenDetectorMLP(nn.Module):
82
+ """input_size = 4 * hidden dim del backbone: (CLS + patch-mean) de la
83
+ rama local concatenado con (CLS + patch-mean) de la rama global."""
84
+
85
+ def __init__(self, input_size: int = 3072, hidden_size: int = 256,
86
+ num_classes: int = 2, dropout: float = 0.3):
87
+ super().__init__()
88
+ self.mlp = nn.Sequential(
89
+ nn.Linear(input_size, hidden_size),
90
+ nn.GELU(),
91
+ nn.BatchNorm1d(hidden_size),
92
+ nn.Dropout(dropout),
93
+ nn.Linear(hidden_size, hidden_size // 2),
94
+ nn.GELU(),
95
+ nn.Dropout(dropout),
96
+ nn.Linear(hidden_size // 2, num_classes),
97
+ )
98
+
99
+ def forward(self, x):
100
+ return self.mlp(x)
101
+
102
+
103
+ def get_num_register_tokens(backbone) -> int:
104
+ return getattr(backbone.config, "num_register_tokens", 0)
105
+
106
+
107
+ @torch.no_grad()
108
+ def extract_dual_features(backbone, local_images, global_images):
109
+ """Misma lógica que en el entrenamiento pero solo-inferencia: concatena
110
+ local+global en el batch, un único forward, CLS + promedio de patch
111
+ tokens (sin register tokens) de cada rama."""
112
+ n_reg = get_num_register_tokens(backbone)
113
+ batch = torch.cat([local_images, global_images], dim=0)
114
+
115
+ if device.type == "cuda":
116
+ with torch.autocast(device_type="cuda", dtype=torch.float16):
117
+ out = backbone(pixel_values=batch)
118
+ else:
119
+ out = backbone(pixel_values=batch)
120
+ hidden = out.last_hidden_state.float()
121
+
122
+ cls_tok = hidden[:, 0, :]
123
+ patch_mean = hidden[:, 1 + n_reg:, :].mean(dim=1)
124
+ feat = torch.cat([cls_tok, patch_mean], dim=-1)
125
+
126
+ B = local_images.size(0)
127
+ local_feat, global_feat = feat[:B], feat[B:]
128
+ return torch.cat([local_feat, global_feat], dim=-1) # (B, 4*hidden)
129
+
130
+
131
+ def load_classes() -> list:
132
+ if not os.path.exists(CFG["classes_path"]):
133
+ raise FileNotFoundError(
134
+ f"No encuentro {CFG['classes_path']}. Corre primero el script de entrenamiento."
135
+ )
136
+ with open(CFG["classes_path"]) as f:
137
+ return json.load(f)
138
+
139
+
140
+ def load_models(num_classes: int):
141
+ print(f"Cargando backbone {CFG['backbone_name']}...")
142
+ backbone = AutoModel.from_pretrained(CFG["backbone_name"]).to(device)
143
+ backbone.eval()
144
+ for p in backbone.parameters():
145
+ p.requires_grad_(False)
146
+
147
+ # Si el entrenamiento descongeló los últimos N bloques, esos bloques
148
+ # tienen pesos fine-tuneados guardados aparte -- hay que cargarlos, si
149
+ # no, estaríamos evaluando con el backbone pre-entrenado original y
150
+ # los resultados no coincidirían con el checkpoint de la cabeza.
151
+ n_unfreeze = CFG["unfreeze_last_n_blocks"]
152
+ if n_unfreeze > 0:
153
+ if os.path.exists(CFG["backbone_ckpt_path"]):
154
+ total_layers = len(backbone.encoder.layer)
155
+ unfrozen_state = torch.load(CFG["backbone_ckpt_path"], map_location=device)
156
+ for i, layer in enumerate(backbone.encoder.layer[total_layers - n_unfreeze:]):
157
+ layer.load_state_dict(unfrozen_state[f"layer.{total_layers - n_unfreeze + i}"])
158
+ print(f"Pesos fine-tuneados de los últimos {n_unfreeze} bloques cargados "
159
+ f"desde {CFG['backbone_ckpt_path']}")
160
+ else:
161
+ print(f"AVISO: unfreeze_last_n_blocks={n_unfreeze} pero no existe "
162
+ f"{CFG['backbone_ckpt_path']}. Evaluando con el backbone SIN fine-tunear "
163
+ f"-- los resultados pueden no coincidir con el val_acc del entrenamiento.")
164
+
165
+ feat_dim = backbone.config.hidden_size * 4
166
+ head = ScreenDetectorMLP(input_size=feat_dim, hidden_size=CFG["hidden_size"],
167
+ num_classes=num_classes, dropout=CFG["dropout"]).to(device)
168
+ head.load_state_dict(torch.load(CFG["ckpt_path"], map_location=device))
169
+ head.eval()
170
+ print(f"Pesos del MLP cargados desde {CFG['ckpt_path']}")
171
+ return backbone, head
172
+
173
+
174
+ # ---------------------------------------------------------------------------
175
+ # Dataset de evaluación simple (una vista local + una vista global por imagen)
176
+ # ---------------------------------------------------------------------------
177
+ def _index_samples(root_dir):
178
+ samples = []
179
+ for dirpath, _, filenames in os.walk(root_dir):
180
+ for fname in filenames:
181
+ lower = fname.lower()
182
+ if not lower.endswith((".jpg", ".jpeg", ".png")):
183
+ continue
184
+ if "_gt" in lower:
185
+ label = 0
186
+ elif "_moire" in lower:
187
+ label = 1
188
+ else:
189
+ continue
190
+ samples.append((os.path.join(dirpath, fname), label))
191
+ if not samples:
192
+ raise RuntimeError(f"No se encontraron imágenes '_gt'/'_moire' en {root_dir}")
193
+ return samples
194
+
195
+
196
+ class MoireDualDataset(Dataset):
197
+ """Igual que en el entrenamiento: cada muestra devuelve (local, global,
198
+ label), con transforms deterministas (sin augmentación) para evaluación."""
199
+
200
+ def __init__(self, root_dir):
201
+ self.samples = _index_samples(root_dir)
202
+
203
+ def __len__(self):
204
+ return len(self.samples)
205
+
206
+ def __getitem__(self, idx):
207
+ path, label = self.samples[idx]
208
+ img = Image.open(path).convert("RGB")
209
+ local_img = local_eval_transform(img)
210
+ global_img = global_eval_transform(img)
211
+ return local_img, global_img, label
212
+
213
+
214
+ # ---------------------------------------------------------------------------
215
+ # Multi-crop TTA (opcional, --tta): mismo criterio que evaluate_tta() en el
216
+ # script de entrenamiento -- centro + 4 esquinas como crops locales, más la
217
+ # vista global, promediando las probabilidades.
218
+ # ---------------------------------------------------------------------------
219
+ def make_local_crops(img: Image.Image, n_crops: int, crop_size: int) -> torch.Tensor:
220
+ w, h = img.size
221
+ cs = crop_size
222
+ cx, cy = max((w - cs) // 2, 0), max((h - cs) // 2, 0)
223
+ positions = [(cx, cy),
224
+ (0, 0), (max(w - cs, 0), 0), (0, max(h - cs, 0)), (max(w - cs, 0), max(h - cs, 0))]
225
+ while len(positions) < n_crops:
226
+ positions.append((random.randint(0, max(w - cs, 0)), random.randint(0, max(h - cs, 0))))
227
+ positions = positions[:n_crops]
228
+
229
+ crops = []
230
+ for x, y in positions:
231
+ crop = img.crop((x, y, x + cs, y + cs))
232
+ if crop.size != (cs, cs):
233
+ crop = crop.resize((cs, cs))
234
+ crops.append(tta_base_transform(crop))
235
+ return torch.stack(crops) # (n_crops, C, H, W)
236
+
237
+
238
+ def make_global_view(img: Image.Image, global_resize: int, crop_size: int) -> torch.Tensor:
239
+ g = TF.resize(img, [global_resize, global_resize])
240
+ g = TF.center_crop(g, [crop_size, crop_size])
241
+ return tta_base_transform(g)
242
+
243
+
244
+ @torch.no_grad()
245
+ def predict_image_tta(backbone, head, img: Image.Image):
246
+ """Predicción multi-crop para UNA imagen ya abierta (PIL). Devuelve
247
+ (pred_idx, confidence, probs_tensor)."""
248
+ local_crops = make_local_crops(img, CFG["tta_crops"], CFG["input_size"]).to(device)
249
+ global_view = make_global_view(img, CFG["global_resize"], CFG["input_size"])
250
+ global_crops = global_view.unsqueeze(0).expand(CFG["tta_crops"], -1, -1, -1).contiguous().to(device)
251
+
252
+ feats = extract_dual_features(backbone, local_crops, global_crops)
253
+ logits = head(feats)
254
+ probs = torch.softmax(logits, dim=-1).mean(dim=0)
255
+ conf, pred = probs.max(dim=0)
256
+ return pred.item(), conf.item(), probs.cpu()
257
+
258
+
259
+ # ---------------------------------------------------------------------------
260
+ # 1. Evaluación sobre el test set original
261
+ # ---------------------------------------------------------------------------
262
+ def evaluate_test_dataset(backbone, head, classes: list, use_tta: bool):
263
+ import kagglehub
264
+
265
+ dataset_root = kagglehub.dataset_download(CFG["dataset_slug"]) # cacheado
266
+ test_dir = os.path.join(dataset_root, "test", "test")
267
+ if not os.path.isdir(test_dir):
268
+ print(f"Aviso: no encuentro {test_dir}.")
269
+ return
270
+
271
+ num_classes = len(classes)
272
+ confusion = torch.zeros(num_classes, num_classes, dtype=torch.long)
273
+ correct, total = 0, 0
274
+
275
+ if use_tta:
276
+ samples = _index_samples(test_dir)
277
+ print(f"\nEvaluando (TTA, {CFG['tta_crops']} crops/imagen) sobre "
278
+ f"{len(samples)} imágenes del test set...")
279
+ for path, label in samples:
280
+ img = Image.open(path).convert("RGB")
281
+ pred, _, _ = predict_image_tta(backbone, head, img)
282
+ confusion[label, pred] += 1
283
+ correct += int(pred == label)
284
+ total += 1
285
+ else:
286
+ test_ds = MoireDualDataset(test_dir)
287
+ loader = DataLoader(test_ds, batch_size=CFG["batch_size"], shuffle=False,
288
+ num_workers=min(4, os.cpu_count() or 1),
289
+ pin_memory=device.type == "cuda")
290
+ print(f"\nEvaluando sobre {len(test_ds)} imágenes del test set...")
291
+ for local_imgs, global_imgs, labels in loader:
292
+ local_imgs, global_imgs = local_imgs.to(device), global_imgs.to(device)
293
+ feats = extract_dual_features(backbone, local_imgs, global_imgs)
294
+ logits = head(feats)
295
+ pred = logits.argmax(dim=1).cpu()
296
+ for t, p in zip(labels, pred):
297
+ confusion[t, p] += 1
298
+ correct += (pred == labels).sum().item()
299
+ total += labels.size(0)
300
+
301
+ acc = 100 * correct / total
302
+ print(f"\nAccuracy en test set: {acc:.2f}% ({correct}/{total})")
303
+
304
+ print("\nMatriz de confusión (filas = real, columnas = predicho):")
305
+ header = "".join(f"{c[:10]:>12}" for c in classes)
306
+ print(" " * 12 + header)
307
+ for i, row in enumerate(confusion):
308
+ row_str = "".join(f"{v.item():>12}" for v in row)
309
+ print(f"{classes[i][:10]:>12}{row_str}")
310
+
311
+ print("\nPor clase:")
312
+ for i, cls in enumerate(classes):
313
+ tp = confusion[i, i].item()
314
+ fn = confusion[i, :].sum().item() - tp
315
+ fp = confusion[:, i].sum().item() - tp
316
+ precision = tp / (tp + fp) if (tp + fp) else 0.0
317
+ recall = tp / (tp + fn) if (tp + fn) else 0.0
318
+ print(f" {cls}: precision={precision:.3f} recall={recall:.3f} (n={confusion[i, :].sum().item()})")
319
+
320
+
321
+ # ---------------------------------------------------------------------------
322
+ # 2. Predicción sobre fotos propias (sin etiquetas)
323
+ # ---------------------------------------------------------------------------
324
+ IMG_EXTENSIONS = (".jpg", ".jpeg", ".png", ".bmp", ".webp")
325
+
326
+
327
+ def collect_image_paths(path: str) -> list:
328
+ if os.path.isfile(path):
329
+ return [path]
330
+ paths = []
331
+ for ext in IMG_EXTENSIONS:
332
+ paths.extend(glob.glob(os.path.join(path, f"*{ext}")))
333
+ paths.extend(glob.glob(os.path.join(path, f"*{ext.upper()}")))
334
+ return sorted(paths)
335
+
336
+
337
+ @torch.no_grad()
338
+ def predict_own_images(backbone, head, classes: list, images_path: str, use_tta: bool):
339
+ paths = collect_image_paths(images_path)
340
+ if not paths:
341
+ print(f"No encontré imágenes en {images_path}")
342
+ return
343
+
344
+ tag = " (TTA)" if use_tta else ""
345
+ print(f"\nPrediciendo{tag} sobre {len(paths)} imagen(es) propias...")
346
+ for path in paths:
347
+ try:
348
+ img = Image.open(path).convert("RGB")
349
+ except Exception as e:
350
+ print(f" {os.path.basename(path)}: no se pudo abrir ({e})")
351
+ continue
352
+
353
+ if use_tta:
354
+ pred_idx, conf, probs = predict_image_tta(backbone, head, img)
355
+ else:
356
+ local_img = local_eval_transform(img).unsqueeze(0).to(device)
357
+ global_img = global_eval_transform(img).unsqueeze(0).to(device)
358
+ feats = extract_dual_features(backbone, local_img, global_img)
359
+ logits = head(feats)
360
+ probs = torch.softmax(logits, dim=1)[0].cpu()
361
+ conf, pred_idx = probs.max(dim=0)
362
+ conf, pred_idx = conf.item(), pred_idx.item()
363
+
364
+ pred_class = classes[pred_idx]
365
+ print(f" {os.path.basename(path):<40} -> {pred_class:<15} "
366
+ f"(confianza {conf*100:.1f}%) "
367
+ f"[{', '.join(f'{c}={p*100:.1f}%' for c, p in zip(classes, probs.tolist()))}]")
368
+
369
+
370
+ # ---------------------------------------------------------------------------
371
+ # Main
372
+ # ---------------------------------------------------------------------------
373
+ def main():
374
+ parser = argparse.ArgumentParser(description="Evalúa el detector de moiré/pantalla (dual-branch).")
375
+ parser.add_argument("--images", type=str, default=None,
376
+ help="Carpeta (o archivo) con tus propias fotos a evaluar.")
377
+ parser.add_argument("--test-dataset", action="store_true",
378
+ help="Evalúa sobre el split de test del dataset original, con métricas.")
379
+ parser.add_argument("--tta", action="store_true",
380
+ help="Usa multi-crop test-time augmentation (más lento, más preciso).")
381
+ args = parser.parse_args()
382
+
383
+ if not args.images and not args.test_dataset:
384
+ parser.error("Especifica --images, --test-dataset, o ambos.")
385
+
386
+ classes = load_classes()
387
+ backbone, head = load_models(num_classes=len(classes))
388
+
389
+ if args.images:
390
+ predict_own_images(backbone, head, classes, args.images, use_tta=args.tta)
391
+
392
+ if args.test_dataset:
393
+ evaluate_test_dataset(backbone, head, classes, use_tta=args.tta)
394
+
395
+
396
+ if __name__ == "__main__":
397
+ main()