hardiksharma6555's picture
Add app.py
48e214f verified
Raw
History Blame Contribute Delete
4.58 kB
import gradio as gr
import torch
import torch.nn as nn
from torchvision import transforms
from transformers import DeiTModel
from PIL import Image
import os
# -----------------------------------------------------------------------------
# Model Definition (re-defined to be self-contained)
# -----------------------------------------------------------------------------
class FingerprintLivenessModel(nn.Module):
"""
Simplified model for fingerprint liveness detection.
DeiT-Tiny backbone (192-dim) -> Fingerprint Expert (192->512) -> Classifier (512->2)
"""
def __init__(self):
super(FingerprintLivenessModel, self).__init__()
# Load pre-trained DeiT-Tiny model
self.base_model = DeiTModel.from_pretrained(
"facebook/deit-tiny-distilled-patch16-224"
)
# Fingerprint-specific expert layer
self.fingerprint_expert = nn.Linear(192, 512)
# Final classifier (2 classes: spoof=0, live=1)
self.classifier = nn.Linear(512, 2)
def forward(self, pixel_values):
outputs = self.base_model(pixel_values)
cls_embeddings = outputs.last_hidden_state[:, 0, :]
expert_features = self.fingerprint_expert(cls_embeddings)
logits = self.classifier(expert_features)
return logits
# -----------------------------------------------------------------------------
# Global Variables & Initialization
# -----------------------------------------------------------------------------
DEVICE = "cpu" # Force CPU for Hugging Face Spaces free tier stability
MODEL_PATH = "model.pth"
def load_model():
if not os.path.exists(MODEL_PATH):
raise FileNotFoundError(f"Model weights not found at: {MODEL_PATH}")
print(f"Loading model from {MODEL_PATH}...")
model = FingerprintLivenessModel()
# Load weights
checkpoint = torch.load(MODEL_PATH, map_location=DEVICE)
model_state_dict = model.state_dict()
# Filter and load weights
filtered_checkpoint = {}
for key in checkpoint.keys():
if key.startswith('base_model') or key.startswith('fingerprint_expert') or key.startswith('classifier'):
if key in model_state_dict:
filtered_checkpoint[key] = checkpoint[key]
model.load_state_dict(filtered_checkpoint, strict=False)
model.to(DEVICE)
model.eval()
return model
# Initialize model once
try:
model = load_model()
print("Model loaded successfully!")
except Exception as e:
print(f"Error loading model: {e}")
model = None
# -----------------------------------------------------------------------------
# Preprocessing
# -----------------------------------------------------------------------------
def preprocess_image(image):
normalize = transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
normalize
])
return transform(image).unsqueeze(0)
# -----------------------------------------------------------------------------
# Prediction Function
# -----------------------------------------------------------------------------
def predict(image):
if model is None:
return "Model not loaded"
if image is None:
return "Please upload an image."
try:
# Preprocess
input_tensor = preprocess_image(image).to(DEVICE)
# Inference
with torch.no_grad():
logits = model(input_tensor)
probs = torch.softmax(logits, dim=1)
spoof_prob = probs[0, 0].item()
live_prob = probs[0, 1].item()
return {
"Live": live_prob,
"Spoof": spoof_prob
}
except Exception as e:
return f"Error during prediction: {str(e)}"
# -----------------------------------------------------------------------------
# Gradio Interface
# -----------------------------------------------------------------------------
title = "Fingerprint Liveness Detection"
description = """
Upload a fingerprint image to check if it's **Live** or **Spoof**.
This model uses a DeiT-Tiny transformer backbone with a fingerprint-specific expert layer.
"""
iface = gr.Interface(
fn=predict,
inputs=gr.Image(type="pil", label="Fingerprint Image"),
outputs=gr.Label(num_top_classes=2, label="Prediction"),
title=title,
description=description,
examples=[]
)
if __name__ == "__main__":
iface.launch()