UserPollo's picture JvL359's picture
Update critical error of README.md (#1)
b6e5890
|
Raw
History Blame Contribute Delete
5.18 kB
---
license: apache-2.0
tags:
- vision
- image-classification
- dinov2
- moire-detection
- screen-detection
- anti-spoofing
- digital-forensics
datasets:
- soumikrakshit/uhdm-dataset
pipeline_tag: image-classification
base_model:
- facebook/dinov2-with-registers-base
---
# Dual-Branch DINOv2 Moiré & Screen Recapture Detector
A robust digital forensics classifier built on **DINOv2 (with registers)** designed to detect screen recaptures and Moiré patterns.
Optimized for anti-spoofing pipelines and automated asset valuation platforms, this model overcomes the traditional scale dilemma in Moiré detection by combining an un-resized native crop (capturing high-frequency pixel interference) with a global thumbnail (capturing screen-wide periodic banding).
---
## 📁 Repository Files
* **`best_screen_detector_backbone.pt`**: Weights for the fine-tuned top transformer blocks of the DINOv2 backbone.
* **`best_screen_detector_mlp.pt`**: Weights for the 2-layer classification MLP head.
* **`classes.json`**: Class index mapping (`0: "gt"`, `1: "moire"`).
---
## 🚀 Quick Start & Inference
### 1. Requirements
If you are running this in **Google Colab**, you do not need to run `pip install` for most of these packages, as PyTorch, Transformers, and Pillow are pre-installed. You only need to ensure `huggingface-hub` is up to date.
For local environments, install the dependencies:
```bash
pip install torch torchvision transformers pillow huggingface-hub requests
```
### 2. Inference Script
```python
import torch
import torch.nn as nn
import json
import requests
from PIL import Image
from torchvision import transforms
from transformers import AutoModel
from huggingface_hub import hf_hub_download
# 1. Download weights and classes from the Hub
repo_id = "UserPollo/moire-pattern-detector"
backbone_ckpt = hf_hub_download(repo_id=repo_id, filename="best_screen_detector_backbone.pt")
mlp_ckpt = hf_hub_download(repo_id=repo_id, filename="best_screen_detector_mlp.pt")
classes_file = hf_hub_download(repo_id=repo_id, filename="classes.json")
with open(classes_file, "r") as f:
classes = json.load(f)
# 2. Define the custom MLP Head
class ScreenDetectorMLP(nn.Module):
def __init__(self, input_size=3072, hidden_size=256, num_classes=2, dropout=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)
# 3. Load Models and apply weights
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
backbone = AutoModel.from_pretrained("facebook/dinov2-with-registers-base").to(device)
backbone.eval()
for p in backbone.parameters():
p.requires_grad_(False)
# Load fine-tuned weights into the last 2 blocks of the backbone
total_layers = len(backbone.encoder.layer)
unfrozen_state = torch.load(backbone_ckpt, map_location=device, weights_only=True)
for i, layer in enumerate(backbone.encoder.layer[total_layers - 2:]):
layer.load_state_dict(unfrozen_state[f"layer.{total_layers - 2 + i}"])
head = ScreenDetectorMLP(input_size=3072).to(device)
head.load_state_dict(torch.load(mlp_ckpt, map_location=device, weights_only=True))
head.eval()
# 4. Prepare Dual-Branch Image Transforms
IMAGENET_MEAN, IMAGENET_STD = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]
local_transform = transforms.Compose([
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
])
global_transform = transforms.Compose([
transforms.Resize((256, 256)),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
])
# 5. Load Image
url = "https://encrypted-tbn0.gstatic.com/images?q=tbn:ANd9GcS2pnVr5QT5OjVf8H4YtrJfRoPvGBFAG5pBG8F-LfPGB0sEcyF1JkT6Okly&s=10"
img = Image.open(requests.get(url, stream=True).raw).convert("RGB")
local_tensor = local_transform(img).unsqueeze(0).to(device)
global_tensor = global_transform(img).unsqueeze(0).to(device)
# Concatenate for a single forward pass
batch = torch.cat([local_tensor, global_tensor], dim=0)
# 6. Run Inference
with torch.no_grad():
out = backbone(pixel_values=batch)
hidden = out.last_hidden_state.float()
# Extract CLS and patch mean (ignoring register tokens)
n_reg = getattr(backbone.config, "num_register_tokens", 0)
cls_tok = hidden[:, 0, :]
patch_mean = hidden[:, 1 + n_reg:, :].mean(dim=1)
feat = torch.cat([cls_tok, patch_mean], dim=-1)
# Split back into local and global, then concatenate horizontally
local_feat, global_feat = feat[0:1], feat[1:2]
combined_feat = torch.cat([local_feat, global_feat], dim=-1)
# Classify
logits = head(combined_feat)
probs = torch.softmax(logits, dim=1)
conf, pred = probs.max(dim=1)
print(f"Prediction: {classes[int(pred.item())]} (Confidence: {conf.item()*100:.1f}%)")
```