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