ethix's picture
fix: use CLIPImageProcessor and remove direct 384x384 size override to fix prediction discrepancy
93d3c44
Raw History Blame Contribute Delete
3.13 kB
import os
import time
import threading
import torch
import torch.nn as nn
from PIL import Image
from transformers import ViTForImageClassification, CLIPImageProcessor
MODEL_REPO = "buildborderless/CommunityForensics-DeepfakeDet-ViT"
MODEL_TTL_SECONDS = 45 * 60 # 45 minutes
_model_lock = threading.Lock()
_model_instance = None
_model_loaded_at = 0.0
class ViTWrapper(nn.Module):
"""Wrap HF ViT to work with Captum and PyTorch XAI methods."""
def __init__(self, model_name=MODEL_REPO):
super().__init__()
self.model_name = model_name
self.model = ViTForImageClassification.from_pretrained(model_name)
self.model.set_attn_implementation('eager') # Required for output_attentions
self.model.eval()
self.processor = CLIPImageProcessor.from_pretrained(model_name)
def forward(self, x):
"""Forward pass taking preprocessed tensor (B, C, H, W) and returning logits (B, 1)."""
outputs = self.model(pixel_values=x)
return outputs.logits
def preprocess(self, image: Image.Image) -> torch.Tensor:
"""Preprocess PIL image to model input tensor matching preprocessor_config (shortest_edge=440, center_crop=384)."""
if image.mode != "RGB":
image = image.convert("RGB")
inputs = self.processor(
images=image,
return_tensors="pt",
)
return inputs.pixel_values # (1, 3, 384, 384)
def predict(self, image: Image.Image) -> dict:
"""Get prediction label, probability, and raw logit."""
tensor = self.preprocess(image)
device = next(self.model.parameters()).device
tensor = tensor.to(device)
with torch.no_grad():
logit = self.forward(tensor).item()
prob = float(torch.sigmoid(torch.tensor(logit)).item())
pred = "FAKE" if prob > 0.5 else "REAL"
conf = prob if pred == "FAKE" else (1.0 - prob)
return {
"prediction": pred,
"probability": prob,
"confidence_pct": round(conf * 100, 2),
"logit": round(logit, 4),
}
def get_model(device: str = None) -> ViTWrapper:
"""Thread-safe singleton model loader with periodic background refresh."""
global _model_instance, _model_loaded_at
with _model_lock:
now = time.time()
if _model_instance is None or (now - _model_loaded_at) > MODEL_TTL_SECONDS:
# Maximize CPU utilization
torch.set_num_threads(os.cpu_count() or 4)
torch.set_num_interop_threads(max(1, (os.cpu_count() or 4) // 2))
if device is None:
device = "cpu" # Default to CPU for CPU-based HF Space & low VRAM stability
wrapper = ViTWrapper(MODEL_REPO)
try:
wrapper.to(device)
except Exception as e:
print(f"[model] Target device {device} failed ({e}), falling back to CPU.")
wrapper.to("cpu")
_model_instance = wrapper
_model_loaded_at = now
return _model_instance