aaryank1411 commited on
Commit
b23e468
·
verified ·
1 Parent(s): fdcd65f

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +76 -30
app.py CHANGED
@@ -1,67 +1,113 @@
1
- # Cell 3: Twin CNN Architecture and Contrastive Loss
 
2
  import torch.nn as nn
3
  import torch.nn.functional as F
 
 
 
4
 
 
 
 
 
5
  class SigNetTwinEncoder(nn.Module):
6
  def __init__(self):
7
  super(SigNetTwinEncoder, self).__init__()
8
-
9
  self.conv = nn.Sequential(
10
- # Input: (1, 220, 155)
11
  nn.Conv2d(1, 32, kernel_size=5, stride=1, padding=2),
12
  nn.BatchNorm2d(32),
13
  nn.ReLU(inplace=True),
14
- nn.MaxPool2d(2, 2), # -> (32, 110, 77)
15
 
16
  nn.Conv2d(32, 64, kernel_size=5, stride=1, padding=2),
17
  nn.BatchNorm2d(64),
18
  nn.ReLU(inplace=True),
19
- nn.MaxPool2d(2, 2), # -> (64, 55, 38)
20
 
21
  nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
22
  nn.BatchNorm2d(128),
23
  nn.ReLU(inplace=True),
24
- nn.MaxPool2d(2, 2) # -> (128, 27, 19)
25
  )
26
-
27
- # Flattened feature map: 128 * 27 * 19 = 65,664
28
  self.fc = nn.Sequential(
29
  nn.Linear(128 * 27 * 19, 256),
30
  nn.ReLU(inplace=True),
31
  nn.Dropout(0.3),
32
- nn.Linear(256, 128) # 128-dimensional embedding
33
  )
34
 
35
  def forward_once(self, x):
36
  features = self.conv(x)
37
  features = features.view(features.size(0), -1)
38
  embeddings = self.fc(features)
39
- # L2-normalization on unit hypersphere
40
- embeddings = F.normalize(embeddings, p=2, dim=1)
41
- return embeddings
42
 
43
  def forward(self, input1, input2):
44
  out1 = self.forward_once(input1)
45
  out2 = self.forward_once(input2)
46
  return out1, out2
47
 
48
- class ContrastiveLoss(nn.Module):
49
- def __init__(self, margin=1.0):
50
- super(ContrastiveLoss, self).__init__()
51
- self.margin = margin
52
 
53
- def forward(self, out1, out2, label):
54
- # Euclidean distance
55
- euclidean_distance = F.pairwise_distance(out1, out2)
56
- # label=1 (genuine): minimize distance^2
57
- # label=0 (forged): push distance beyond margin
58
- loss_contrastive = torch.mean(
59
- label * torch.pow(euclidean_distance, 2) +
60
- (1.0 - label) * torch.pow(torch.clamp(self.margin - euclidean_distance, min=0.0), 2)
61
- )
62
- return loss_contrastive
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
 
64
- # Verify parameter count
65
- model_check = SigNetTwinEncoder()
66
- param_count = sum(p.numel() for p in model_check.parameters() if p.requires_grad)
67
- print(f"Total Trainable Parameters: {param_count:,} (~{param_count / 1e6:.2f}M)")
 
1
+ import os
2
+ import torch
3
  import torch.nn as nn
4
  import torch.nn.functional as F
5
+ import cv2
6
+ import numpy as np
7
+ import gradio as gr
8
 
9
+ # Force CPU inference for Hugging Face free tier
10
+ device = torch.device("cpu")
11
+
12
+ # 1. Exact Architecture matching the trained weights
13
  class SigNetTwinEncoder(nn.Module):
14
  def __init__(self):
15
  super(SigNetTwinEncoder, self).__init__()
 
16
  self.conv = nn.Sequential(
 
17
  nn.Conv2d(1, 32, kernel_size=5, stride=1, padding=2),
18
  nn.BatchNorm2d(32),
19
  nn.ReLU(inplace=True),
20
+ nn.MaxPool2d(2, 2),
21
 
22
  nn.Conv2d(32, 64, kernel_size=5, stride=1, padding=2),
23
  nn.BatchNorm2d(64),
24
  nn.ReLU(inplace=True),
25
+ nn.MaxPool2d(2, 2),
26
 
27
  nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1),
28
  nn.BatchNorm2d(128),
29
  nn.ReLU(inplace=True),
30
+ nn.MaxPool2d(2, 2)
31
  )
 
 
32
  self.fc = nn.Sequential(
33
  nn.Linear(128 * 27 * 19, 256),
34
  nn.ReLU(inplace=True),
35
  nn.Dropout(0.3),
36
+ nn.Linear(256, 128)
37
  )
38
 
39
  def forward_once(self, x):
40
  features = self.conv(x)
41
  features = features.view(features.size(0), -1)
42
  embeddings = self.fc(features)
43
+ return F.normalize(embeddings, p=2, dim=1)
 
 
44
 
45
  def forward(self, input1, input2):
46
  out1 = self.forward_once(input1)
47
  out2 = self.forward_once(input2)
48
  return out1, out2
49
 
50
+ # 2. Instantiate and load model weights safely
51
+ model = SigNetTwinEncoder().to(device)
52
+ model_path = "signet_twin_model.pth"
 
53
 
54
+ if os.path.exists(model_path):
55
+ state_dict = torch.load(model_path, map_location=device)
56
+ model.load_state_dict(state_dict)
57
+ model.eval()
58
+ print("Model weights loaded successfully.")
59
+ else:
60
+ print(f"Warning: {model_path} not found in the repository root.")
61
+
62
+ OPTIMAL_THRESHOLD = 0.5581
63
+
64
+ # 3. Preprocessing function
65
+ def preprocess_for_inference(image):
66
+ if image is None:
67
+ return None
68
+ if len(image.shape) == 3:
69
+ gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
70
+ else:
71
+ gray = image
72
+
73
+ _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)
74
+ resized = cv2.resize(binary, (155, 220))
75
+ normalized = resized.astype(np.float32) / 255.0
76
+ tensor = torch.tensor(normalized, dtype=torch.float32).unsqueeze(0).unsqueeze(0)
77
+ return tensor.to(device)
78
+
79
+ # 4. Verification Logic
80
+ def verify_signatures(ref_img, query_img):
81
+ if ref_img is None or query_img is None:
82
+ return "Please upload both reference and query signatures."
83
+
84
+ t_ref = preprocess_for_inference(ref_img)
85
+ t_query = preprocess_for_inference(query_img)
86
+
87
+ with torch.no_grad():
88
+ out_ref, out_query = model(t_ref, t_query)
89
+ distance = F.pairwise_distance(out_ref, out_query).item()
90
+
91
+ if distance <= OPTIMAL_THRESHOLD:
92
+ verdict = "GENUINE MATCH"
93
+ symbol = "✅"
94
+ else:
95
+ verdict = "FORGERY DETECTED"
96
+ symbol = "❌"
97
+
98
+ return f"{symbol} {verdict}\n\n• Euclidean Distance: {distance:.4f}\n• Decision Threshold: {OPTIMAL_THRESHOLD:.4f}"
99
+
100
+ # 5. Gradio Interface
101
+ interface = gr.Interface(
102
+ fn=verify_signatures,
103
+ inputs=[
104
+ gr.Image(label="Reference Signature (Known Genuine)"),
105
+ gr.Image(label="Query Signature (To be Verified)")
106
+ ],
107
+ outputs=gr.Textbox(label="Biometric Verification Verdict"),
108
+ title="SigNet-Verify: Offline Signature Forgery Detection",
109
+ description="Siamese Neural Network (Twin CNN) for writer-independent biometric signature verification."
110
+ )
111
 
112
+ if __name__ == "__main__":
113
+ interface.launch()