Chethanand's picture
Upload 24 files
eb3afa1 verified
Raw
History Blame Contribute Delete
2.18 kB
import torch
import torch.nn as nn
import torchvision.models as models
class ResNetLSTM(nn.Module):
def __init__(self, num_classes=1, hidden_size=256, num_layers=2):
super(ResNetLSTM, self).__init__()
# Load Pretrained ResNet
resnet = models.resnet18(pretrained=True)
# Remove last fully connected layer
# ResNet18 fc input size is 512
modules = list(resnet.children())[:-1]
self.resnet = nn.Sequential(*modules)
# Freeze ResNet params (Optional: unfreeze later or partial unfreeze)
# For now, let's fine-tune all or freeze?
# Fine-tuning is usually better if we have enough data.
# User has ~4k videos, which is decent. Let's NOT freeze.
self.lstm = nn.LSTM(
input_size=512,
hidden_size=hidden_size,
num_layers=num_layers,
batch_first=True,
dropout=0.5
)
self.fc = nn.Linear(hidden_size, num_classes)
self.sigmoid = nn.Sigmoid() # For binary classification logic if needed manually, but we use BCEWithLogitsLoss
def forward(self, x):
# x shape: (Batch, Frames, Channels, Height, Width)
# Example: (B, 30, 3, 224, 224)
b, t, c, h, w = x.size()
# Flatten batch and frames for CNN processing
# (B*T, 3, 224, 224)
c_in = x.view(b * t, c, h, w)
# CNN Feature Extraction
# Output: (B*T, 512, 1, 1) -> squeeze -> (B*T, 512)
features = self.resnet(c_in)
features = features.view(features.size(0), -1)
# Reshape back to (B, T, Features)
r_in = features.view(b, t, -1)
# LSTM Temporal processing
# Output: (B, T, Hidden)
# hidden/cell: (Layers, B, Hidden)
lstm_out, _ = self.lstm(r_in)
# Take the output from the last time step
# (B, Hidden)
last_out = lstm_out[:, -1, :]
# Classification
out = self.fc(last_out)
return out