Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
[![Tests](https://img.shields.io/badge/Tests-PyTest-green.svg)](tests/)


# Facial Expression Recognition
# EmoteVision

A compact end-to-end training and evaluation pipeline for facial expression classification using a Hugging Face dataset and a ResNet50 backbone.

Expand Down
4 changes: 2 additions & 2 deletions src/models/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
from src.models.facial_recognition_model import FacialRecognitionModel
from src.models.emote_vision_model import EmoteVisionModel

__all__ = ["FacialRecognitionModel"]
__all__ = ["EmoteVisionModel"]
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,13 @@
import torch.nn.functional as F
from torchvision.models import resnet50, ResNet50_Weights

class FacialRecognitionModel(nn.Module):

class EmoteVisionModel(nn.Module):
def __init__(self, embedding_size: int = 512, num_classes: int = 7):
super().__init__()

# Pretrained ResNet50 backbone
self.base_model = resnet50(weights = ResNet50_Weights.DEFAULT)
self.base_model = resnet50(weights=ResNet50_Weights.DEFAULT)

in_features = self.base_model.fc.in_features
self.base_model.fc = nn.Identity()
Expand All @@ -27,12 +28,11 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
Returns logits suitable for `nn.CrossEntropyLoss`.
"""
features = self.base_model(x)
vector = torch.flatten(features, start_dim = 1)
vector = torch.flatten(features, start_dim=1)

embeddings = self.embedding_layer(vector)
embeddings = F.relu(embeddings)

logits = self.classifier(embeddings)

return logits

6 changes: 3 additions & 3 deletions src/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from rich.panel import Panel

from src import DataLoader, Trainer, Evaluator
from src.models import FacialRecognitionModel
from src.models import EmoteVisionModel

console = Console()

Expand All @@ -22,7 +22,7 @@ def run_pipeline(config: dict = None) -> None:
output_dir = config.get("output_dir", "outputs")
os.makedirs(output_dir, exist_ok=True)

console.print(Panel.fit("[bold blue]Facial Expression Recognition Pipeline[/bold blue]"))
console.print(Panel.fit("[bold blue]EmoteVision Pipeline[/bold blue]"))

# Step 1: Data Preparation
with console.status("[bold green]Loading DataModule...[/bold green]", spinner="dots"):
Expand All @@ -32,7 +32,7 @@ def run_pipeline(config: dict = None) -> None:

# Step 2: Model & Optimizer Setup
with console.status("[bold green]Initializing Model & Optimizer...[/bold green]", spinner="dots"):
model = FacialRecognitionModel()
model = EmoteVisionModel()
optimizer = optim.Adam(model.parameters(), lr=config["learning_rate"])
criterion = nn.CrossEntropyLoss()
console.print(" Initialized Model and Optimizer")
Expand Down
6 changes: 3 additions & 3 deletions tests/test_facial_recognition.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,12 @@
import pytest
import torch
from src.models import FacialRecognitionModel
from src.models import EmoteVisionModel

def test_model_output_shape():
batch_size = 4
num_classes = 5

model = FacialRecognitionModel(num_classes = num_classes)
model = EmoteVisionModel(num_classes = num_classes)

fake_images = torch.randn(batch_size, 3, 150, 150)

Expand All @@ -16,7 +16,7 @@ def test_model_output_shape():
assert output.shape == (batch_size, num_classes), f"Expected shape {(batch_size, num_classes)}, got {output.shape}"

def test_backbone_weights_are_frozen():
model = FacialRecognitionModel()
model = EmoteVisionModel()

first_layer_param = next(model.base_model.parameters())

Expand Down
2 changes: 1 addition & 1 deletion tests/test_trainer.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import pytest
from src.models import FacialRecognitionModel
from src.models import EmoteVisionModel
from src import Trainer
import torch
import torch.nn as nn
Expand Down
Loading