| 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 |
|
|
| |
| |
| |
| 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__() |
| |
| self.base_model = DeiTModel.from_pretrained( |
| "facebook/deit-tiny-distilled-patch16-224" |
| ) |
| |
| self.fingerprint_expert = nn.Linear(192, 512) |
| |
| 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 |
|
|
| |
| |
| |
| DEVICE = "cpu" |
| 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() |
| |
| |
| checkpoint = torch.load(MODEL_PATH, map_location=DEVICE) |
| model_state_dict = model.state_dict() |
| |
| |
| 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 |
|
|
| |
| try: |
| model = load_model() |
| print("Model loaded successfully!") |
| except Exception as e: |
| print(f"Error loading model: {e}") |
| model = None |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| def predict(image): |
| if model is None: |
| return "Model not loaded" |
| |
| if image is None: |
| return "Please upload an image." |
| |
| try: |
| |
| input_tensor = preprocess_image(image).to(DEVICE) |
| |
| |
| 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)}" |
|
|
| |
| |
| |
| 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() |
|
|