Server-Main / Version_2 /inference.py
hch-dev
Reorganized and renamed Version folders
e6bbf5e
Raw
History Blame Contribute Delete
2.23 kB
import torch
import timm
import ttach as tta
from PIL import Image
import os
from core.augmentations import get_eval_transforms
class ProductionAnalyzer:
def __init__(self, model_path, num_classes=2):
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.transform = get_eval_transforms()
self.class_names = {0: "Legit", 1: "Phishing"}
print("Loading Production Model...")
base_model = timm.create_model('convnext_tiny', pretrained=False, num_classes=num_classes)
# Load weights and strip the 'module.' prefix caused by AveragedModel saving
state_dict = torch.load(model_path, map_location=self.device, weights_only=True)
clean_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
base_model.load_state_dict(clean_state_dict)
base_model.to(self.device)
base_model.eval()
# Wrap in TTA (Test Time Augmentation)
self.model = tta.ClassificationTTAWrapper(
base_model,
tta.aliases.five_crop_transform(224, 224)
)
def analyze_user_input(self, image_path):
if not os.path.exists(image_path):
return "Error: Image file not found."
image = Image.open(image_path).convert("RGB")
input_tensor = self.transform(image).unsqueeze(0).to(self.device)
with torch.no_grad():
logits = self.model(input_tensor)
probabilities = torch.softmax(logits, dim=1)
confidence, predicted_class = torch.max(probabilities, dim=1)
class_id = predicted_class.item()
return {
"prediction": self.class_names[class_id],
"class_id": class_id,
"confidence_score": f"{confidence.item() * 100:.2f}%"
}
if __name__ == "__main__":
# Ensure production_convnext_ema.pth exists in the directory before running
analyzer = ProductionAnalyzer(model_path="production_convnext_ema.pth")
# Example usage:
# result = analyzer.analyze_user_input("path_to_screenshot_to_test.jpg")
# print(result)