Download xai_engine/model.py from LPX55/DeepfakeDetection-Explainability: direct link, hf CLI and curl.
- Browser
- Download file 3.13 kB
-
https://huggingface.co/spaces/LPX55/DeepfakeDetection-Explainability/resolve/main/xai_engine/model.py
- Command line
-
hf download hf://spaces/LPX55/DeepfakeDetection-Explainability/xai_engine/model.py
-
curl -L -o model.py https://huggingface.co/spaces/LPX55/DeepfakeDetection-Explainability/resolve/main/xai_engine/model.py
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 | |