vaniv commited on
Commit
c3f0b92
·
verified ·
1 Parent(s): 2ce85ab

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +109 -139
app.py CHANGED
@@ -1,55 +1,17 @@
1
- import io, os, numpy as np, gradio as gr
2
- from PIL import Image, ImageChops, ImageDraw
 
 
 
3
  import cv2
4
- from skimage import exposure
 
5
  import mediapipe as mp
6
 
7
- # ====== HF model choice (pick one) ======
8
- HF_MODEL_ID = os.getenv("HF_MODEL_ID", "prithivMLmods/Deep-Fake-Detector-v2-Model") # ViT 224
9
- HF_IMAGE_SIZE = int(os.getenv("HF_IMAGE_SIZE", "224")) # 224 for v2 ViT, 512 for v1 SigLIP
10
-
11
- # ====== HF imports (lazy so app can start even if transformers missing) ======
12
- _hf_loaded = False
13
- _hf_processor = None
14
- _hf_model = None
15
- def _try_load_hf():
16
- global _hf_loaded, _hf_processor, _hf_model
17
- if _hf_loaded:
18
- return True
19
- try:
20
- from transformers import AutoImageProcessor, AutoModelForImageClassification
21
- _hf_processor = AutoImageProcessor.from_pretrained(HF_MODEL_ID)
22
- _hf_model = AutoModelForImageClassification.from_pretrained(HF_MODEL_ID)
23
- _hf_model.eval()
24
- _hf_loaded = True
25
- return True
26
- except Exception as e:
27
- print("HF load failed:", e)
28
- _hf_loaded = False
29
- return False
30
-
31
- def _hf_predict_proba(pil_rgb_face):
32
- """Returns probability that image is deepfake, in [0,1]."""
33
- import torch
34
- with torch.no_grad():
35
- inputs = _hf_processor(images=pil_rgb_face.resize((HF_IMAGE_SIZE, HF_IMAGE_SIZE)), return_tensors="pt")
36
- outputs = _hf_model(**inputs)
37
- logits = outputs.logits[0]
38
- probs = torch.softmax(logits, dim=-1).cpu().numpy()
39
- # Map label -> index; models commonly use ["Deepfake","Realism"] or ["fake","real"]
40
- id2label = _hf_model.config.id2label
41
- lab2idx = {v.lower(): k for k, v in _hf_model.config.label2id.items()}
42
- # Try a few common names
43
- deep_idx = lab2idx.get("deepfake", None)
44
- if deep_idx is None:
45
- deep_idx = lab2idx.get("fake", None)
46
- if deep_idx is None:
47
- # Heuristic: choose the class whose label name contains 'fake'
48
- deep_idx = next((i for i, name in id2label.items() if "fake" in name.lower()), 0)
49
- return float(probs[int(deep_idx)])
50
-
51
- # ====== Face detect / crop (your pipeline) ======
52
- _mp_face = mp.solutions.face_detection.FaceDetection(model_selection=0, min_detection_confidence=0.4)
53
 
54
  def crop_face(pil_img, pad=0.25):
55
  img = np.array(pil_img.convert("RGB"))
@@ -57,7 +19,10 @@ def crop_face(pil_img, pad=0.25):
57
  res = _mp_face.process(cv2.cvtColor(img, cv2.COLOR_RGB2BGR))
58
  if not res.detections:
59
  return pil_img
60
- det = max(res.detections, key=lambda d: d.location_data.relative_bounding_box.width)
 
 
 
61
  b = det.location_data.relative_bounding_box
62
  x, y, bw, bh = b.xmin, b.ymin, b.width, b.height
63
  x1 = int(max(0, (x - pad*bw) * w)); y1 = int(max(0, (y - pad*bh) * h))
@@ -65,73 +30,77 @@ def crop_face(pil_img, pad=0.25):
65
  face = Image.fromarray(img[y1:y2, x1:x2])
66
  return face if face.size[0] > 20 and face.size[1] > 20 else pil_img
67
 
68
- # ====== Heuristic fallback (unchanged core) ======
69
- def _enhance_for_display(pil_img, scale: float):
70
- arr = np.array(pil_img).astype("float32") * scale
71
- arr = np.clip(arr, 0, 255).astype("uint8")
72
- return Image.fromarray(arr)
73
-
74
- def error_level_analysis(pil_img: Image.Image, quality: int = 90):
75
- img = pil_img.convert("RGB")
76
- with io.BytesIO() as buf:
77
- img.save(buf, "JPEG", quality=quality); buf.seek(0)
78
- comp = Image.open(buf).convert("RGB")
79
- diff = ImageChops.difference(img, comp)
80
- extrema = diff.getextrema(); max_diff = max([m for (_, m) in extrema])
81
- scale = 255.0 / max(1, max_diff)
82
- ela_vis = _enhance_for_display(diff, scale)
83
- ela_np = np.array(ela_vis, dtype=np.float32)
84
- mean_intensity = float(ela_np.mean() / 255.0)
85
- return ela_vis, mean_intensity
86
-
87
- def ela_sweep_mean(pil_img, qualities=(95, 90, 85)):
88
- vals = []
89
- for q in qualities:
90
- _, m = error_level_analysis(pil_img, quality=q); vals.append(m)
91
- return float(max(vals)), float(np.mean(vals))
92
-
93
- def fft_high_freq_ratio(pil_img: Image.Image):
94
- y = pil_img.convert("YCbCr").split()[0]
95
- gray = np.array(y, dtype=np.float32)/255.0
96
- h, w = gray.shape
97
- wy, wx = np.hanning(h)[:, None], np.hanning(w)[None, :]
98
- F = np.fft.fftshift(np.fft.fft2(gray * (wy * wx)))
99
- mag = np.log1p(np.abs(F))
100
- cy, cx = h//2, w//2
101
- yy, xx = np.ogrid[:h, :w]; dist = np.sqrt((yy - cy)**2 + (xx - cx)**2)
102
- r_low = min(h, w) * 0.08
103
- low = float(mag[dist <= r_low].sum()); high = float(mag[dist > r_low].sum())
104
- return None, float(high / (high + low + 1e-9))
105
-
106
- def noise_inconsistency(pil_img: Image.Image):
107
- y = pil_img.convert("YCbCr").split()[0]
108
- img = np.array(y, dtype=np.float32)
109
- lap = cv2.Laplacian(img, cv2.CV_32F, ksize=3); lap_abs = np.abs(lap)
110
- tile = 32; H, W = lap_abs.shape; vals = []
111
- for yy in range(0, H, tile):
112
- for xx in range(0, W, tile):
113
- patch = lap_abs[yy:min(yy+tile, H), xx:min(xx+tile, W)]
114
- if patch.size: vals.append(patch.var())
115
- if not vals: return None, 0.0
116
- vals = np.array(vals, dtype=np.float32)
117
- score = float(vals.std() / (vals.mean() + 1e-9))
118
- return None, float(np.tanh(score / 5.0))
119
-
120
- def combine_scores(ela_mean, hf_ratio, noise_incons_score):
121
- w1, w2, w3 = 0.30, 0.40, 0.30
122
- s_ela = np.clip(ela_mean * 3.0, 0, 1)
123
- s_hf = np.clip((hf_ratio - 0.65) / 0.25, 0, 1)
124
- s_noi = np.clip(noise_incons_score, 0, 1)
125
- conf = float(w1*s_ela + w2*s_hf + w3*s_noi)
126
- label = "Likely Manipulated" if conf >= 0.65 else "Likely Authentic"
127
- return label, conf
128
-
129
- # ====== Result card ======
130
- def _result_card(label: str, conf: float, note: str | None = None) -> str:
 
 
 
 
 
131
  pct = max(0.0, min(1.0, conf)) * 100.0
132
  color = "#d84a4a" if label.startswith("Likely Manipulated") else "#2e7d32"
133
  bar_bg = "#e9ecef"
134
- extra = f"<div style='color:#6b7280;font-size:12px;margin-top:10px;text-align:center;'>{note}</div>" if note else ""
135
  return f"""
136
  <div style="max-width:860px;margin:0 auto;">
137
  <div style="border:1px solid #e5e7eb;border-radius:14px;padding:18px 20px;background:#fff;
@@ -144,31 +113,23 @@ def _result_card(label: str, conf: float, note: str | None = None) -> str:
144
  <div style="height:100%;width:{pct:.4f}%;background:{color};"></div>
145
  </div>
146
  </div>
147
- {extra}
148
  </div>
149
  """
150
 
151
- # ====== Inference ======
152
  def analyze(pil_img: Image.Image):
153
  if pil_img is None:
154
  return _result_card("Likely Authentic", 0.0)
155
- face = crop_face(pil_img).convert("RGB")
156
-
157
- if _try_load_hf():
158
- prob_fake = _hf_predict_proba(face)
159
- label = "Likely Manipulated" if prob_fake >= 0.5 else "Likely Authentic"
160
- note = f"HF model: {HF_MODEL_ID}"
161
- return _result_card(label, prob_fake, note=note)
162
-
163
- # Fallback heuristic (if HF model failed)
164
- face = face.resize((512, 512))
165
- _, ela_mean = error_level_analysis(face, quality=90)
166
- _, hf_ratio = fft_high_freq_ratio(face)
167
- _, noi_score = noise_inconsistency(face)
168
- label, conf = combine_scores(ela_mean, hf_ratio, noi_score)
169
- return _result_card(label, conf, note="Heuristic fallback")
170
-
171
- # ====== UI ======
172
  CUSTOM_CSS = """
173
  .gradio-container {max-width: 980px !important;}
174
  .sleek-card {
@@ -176,17 +137,26 @@ CUSTOM_CSS = """
176
  box-shadow: 0 2px 10px rgba(16,24,40,.04); padding: 18px;
177
  }
178
  """
179
- with gr.Blocks(title="Deepfake Detector (Pretrained HF Model)", css=CUSTOM_CSS, theme=gr.themes.Soft()) as demo:
180
- gr.Markdown("<h2 style='text-align:center;margin-bottom:6px;'>Deepfake Detector</h2>"
181
- "<p style='text-align:center;color:#6b7280;'>Face-crop → pretrained classifier → single likelihood.</p>")
 
 
 
182
  with gr.Row():
183
  with gr.Column(scale=6, elem_classes=["sleek-card"]):
184
- inp = gr.Image(type="pil", label="Upload / Paste Image",
185
- sources=["upload", "webcam", "clipboard"],
186
- height=420, show_label=True, interactive=True)
 
 
 
 
 
187
  btn = gr.Button("Analyze", variant="primary", size="lg")
188
  with gr.Column(scale=6):
189
  out = gr.HTML()
 
190
  btn.click(analyze, inputs=inp, outputs=out)
191
  inp.change(analyze, inputs=inp, outputs=out)
192
 
 
1
+ # app.py
2
+ import io
3
+ import numpy as np
4
+ import gradio as gr
5
+ from PIL import Image, ImageDraw
6
  import cv2
7
+ import torch
8
+ from transformers import AutoImageProcessor, ViTForImageClassification
9
  import mediapipe as mp
10
 
11
+ # -------------------- Face crop utilities --------------------
12
+ _mp_face = mp.solutions.face_detection.FaceDetection(
13
+ model_selection=0, min_detection_confidence=0.4
14
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
 
16
  def crop_face(pil_img, pad=0.25):
17
  img = np.array(pil_img.convert("RGB"))
 
19
  res = _mp_face.process(cv2.cvtColor(img, cv2.COLOR_RGB2BGR))
20
  if not res.detections:
21
  return pil_img
22
+ det = max(
23
+ res.detections,
24
+ key=lambda d: d.location_data.relative_bounding_box.width
25
+ )
26
  b = det.location_data.relative_bounding_box
27
  x, y, bw, bh = b.xmin, b.ymin, b.width, b.height
28
  x1 = int(max(0, (x - pad*bw) * w)); y1 = int(max(0, (y - pad*bh) * h))
 
30
  face = Image.fromarray(img[y1:y2, x1:x2])
31
  return face if face.size[0] > 20 and face.size[1] > 20 else pil_img
32
 
33
+ def face_oval_mask(img_pil, shrink=0.80):
34
+ # (Optional) not used by the model; kept if you ever want to mask background
35
+ w, h = img_pil.size
36
+ mask = Image.new("L", (w, h), 0)
37
+ draw = ImageDraw.Draw(mask)
38
+ dx, dy = int((1 - shrink) * w / 2), int((1 - shrink) * h / 2)
39
+ draw.ellipse((dx, dy, w - dx, h - dy), fill=255)
40
+ return np.array(mask, dtype=np.float32) / 255.0
41
+
42
+ # -------------------- HF model: Deepfake vs Realism --------------------
43
+ MODEL_ID = "prithivMLmods/Deep-Fake-Detector-v2-Model"
44
+
45
+ # CPU by default; if you run locally with GPU, you can .to("cuda")
46
+ _hf_processor = AutoImageProcessor.from_pretrained(MODEL_ID)
47
+ _hf_model = ViTForImageClassification.from_pretrained(MODEL_ID)
48
+ _hf_model.eval()
49
+ torch.set_grad_enabled(False)
50
+
51
+ _FAKE_KEYS = ("fake", "deepfake", "manipulated", "spoof", "forged")
52
+
53
+ def _deepfake_index_from_config(cfg) -> int | None:
54
+ """
55
+ Try to find the class index for 'Deepfake' from id2label/label2id.
56
+ This model typically has {0:'Realism', 1:'Deepfake'}.
57
+ """
58
+ # Prefer id2label
59
+ id2label = getattr(cfg, "id2label", None)
60
+ if id2label:
61
+ normalized = {int(k): str(v).lower() for k, v in id2label.items()}
62
+ for idx, lab in normalized.items():
63
+ if any(k in lab for k in _FAKE_KEYS):
64
+ return idx
65
+
66
+ # Fallback to label2id if present
67
+ label2id = getattr(cfg, "label2id", None)
68
+ if label2id:
69
+ inv = {int(v): str(k).lower() for k, v in label2id.items()}
70
+ for idx, lab in inv.items():
71
+ if any(k in lab for k in _FAKE_KEYS):
72
+ return idx
73
+
74
+ return None
75
+
76
+ _DEEP_IDX = _deepfake_index_from_config(_hf_model.config)
77
+
78
+ def _hf_predict_proba(pil_img: Image.Image) -> float:
79
+ """
80
+ Returns P(Deepfake) in [0,1] using the ViT classifier.
81
+ """
82
+ inputs = _hf_processor(images=pil_img.convert("RGB"), return_tensors="pt")
83
+ logits = _hf_model(**inputs).logits # (1, C)
84
+ if logits.shape[-1] == 1:
85
+ # Unlikely for this model, but handle binary-sigmoid heads
86
+ return torch.sigmoid(logits.squeeze(0))[0].item()
87
+
88
+ probs = torch.softmax(logits.squeeze(0), dim=-1).cpu().numpy()
89
+ if _DEEP_IDX is not None and 0 <= _DEEP_IDX < probs.shape[0]:
90
+ return float(probs[_DEEP_IDX])
91
+
92
+ # Binary softmax fallback: assume index 1 = deepfake
93
+ if probs.shape[0] == 2:
94
+ return float(probs[1])
95
+
96
+ # Last resort: take the highest class prob (not ideal, but safe)
97
+ return float(probs.max())
98
+
99
+ # -------------------- Output card --------------------
100
+ def _result_card(label: str, conf: float) -> str:
101
  pct = max(0.0, min(1.0, conf)) * 100.0
102
  color = "#d84a4a" if label.startswith("Likely Manipulated") else "#2e7d32"
103
  bar_bg = "#e9ecef"
 
104
  return f"""
105
  <div style="max-width:860px;margin:0 auto;">
106
  <div style="border:1px solid #e5e7eb;border-radius:14px;padding:18px 20px;background:#fff;
 
113
  <div style="height:100%;width:{pct:.4f}%;background:{color};"></div>
114
  </div>
115
  </div>
 
116
  </div>
117
  """
118
 
119
+ # -------------------- Gradio handler --------------------
120
  def analyze(pil_img: Image.Image):
121
  if pil_img is None:
122
  return _result_card("Likely Authentic", 0.0)
123
+
124
+ # Focus on the face to reduce background false positives
125
+ face = crop_face(pil_img)
126
+ face = face.convert("RGB").resize((224, 224)) # ViT expects 224x224
127
+
128
+ p_fake = _hf_predict_proba(face)
129
+ label = "Likely Manipulated" if p_fake >= 0.65 else "Likely Authentic"
130
+ return _result_card(label, p_fake)
131
+
132
+ # -------------------- UI --------------------
 
 
 
 
 
 
 
133
  CUSTOM_CSS = """
134
  .gradio-container {max-width: 980px !important;}
135
  .sleek-card {
 
137
  box-shadow: 0 2px 10px rgba(16,24,40,.04); padding: 18px;
138
  }
139
  """
140
+
141
+ with gr.Blocks(title="Deepfake Detector (ViT)", css=CUSTOM_CSS, theme=gr.themes.Soft()) as demo:
142
+ gr.Markdown(
143
+ "<h2 style='text-align:center;margin-bottom:6px;'>Deepfake Detector (ViT)</h2>"
144
+ "<p style='text-align:center;color:#6b7280;'>Upload an image to get a single, clean likelihood estimate using a fine-tuned Vision Transformer.</p>"
145
+ )
146
  with gr.Row():
147
  with gr.Column(scale=6, elem_classes=["sleek-card"]):
148
+ inp = gr.Image(
149
+ type="pil",
150
+ label="Upload / Paste Image",
151
+ sources=["upload", "webcam", "clipboard"],
152
+ height=420,
153
+ show_label=True,
154
+ interactive=True,
155
+ )
156
  btn = gr.Button("Analyze", variant="primary", size="lg")
157
  with gr.Column(scale=6):
158
  out = gr.HTML()
159
+
160
  btn.click(analyze, inputs=inp, outputs=out)
161
  inp.change(analyze, inputs=inp, outputs=out)
162