Hackathon4_Face_Expression / app /Hackathon_setup /exp_recognition_model.py
aruntandra's picture
Upload exp_recognition_model.py
0c5cb04 verified
Raw
History Blame Contribute Delete
2.05 kB
import os
import torch
import torch.nn as nn
import torchvision.models as models
from torchvision import transforms
from PIL import Image
# Must match train_dataset.class_to_idx
classes = {
0: "ANGER",
1: "DISGUST",
2: "FEAR",
3: "HAPPINESS",
4: "NEUTRAL",
5: "SADNESS",
6: "SURPRISE",
}
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]
IMG_SIZE = 100
class facExpRec(nn.Module):
"""
ResNet18 backbone with custom fc for 7 expressions.
Deployment-safe: weights=None, we load your checkpoint.
"""
def __init__(self, num_classes: int = 7):
super().__init__()
self.backbone = models.resnet18(weights=None)
in_features = self.backbone.fc.in_features
self.backbone.fc = nn.Linear(in_features, num_classes)
def forward(self, x):
return self.backbone(x)
def ensure_rgb(image: Image.Image) -> Image.Image:
if image.mode != "RGB":
image = image.convert("RGB")
return image
def get_transform(img_size: int = IMG_SIZE):
return transforms.Compose([
transforms.Lambda(ensure_rgb),
transforms.Resize((img_size, img_size)),
transforms.ToTensor(),
transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
])
trnscm = get_transform(IMG_SIZE)
def load_model(checkpoint_path: str, device: str = "cpu", num_classes: int = 7):
model = facExpRec(num_classes=num_classes) # your current class with self.backbone
state = torch.load(checkpoint_path, map_location=device)
# handle wrapper dicts
if isinstance(state, dict) and "state_dict" in state:
state = state["state_dict"]
# If checkpoint keys are like "conv1.weight" but model expects "backbone.conv1.weight"
# then add the "backbone." prefix
sample_key = next(iter(state.keys()))
if not sample_key.startswith("backbone."):
state = {f"backbone.{k}": v for k, v in state.items()}
model.load_state_dict(state, strict=True)
model.to(device)
model.eval()
return model