Deepfake / model.py
Prabhsimran09's picture
Update model.py
b85a843 verified
Raw
History Blame Contribute Delete
3.68 kB
import torch
import torch.nn as nn
class ResNeXtBottleneck(nn.Module):
expansion = 2
def __init__(self, inplanes, planes, stride=1, downsample=None, base_width=2, cardinality=32):
super(ResNeXtBottleneck, self).__init__()
D = int(planes * (base_width / 64.)) * cardinality # = 1024 for planes=1024
C = cardinality
self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)
self.bn1 = nn.BatchNorm2d(planes)
self.conv2 = nn.Conv2d(planes, D, kernel_size=3, stride=stride,
padding=1, groups=C, bias=False)
self.bn2 = nn.BatchNorm2d(D)
self.conv3 = nn.Conv2d(D, planes * self.expansion, kernel_size=1, bias=False)
self.bn3 = nn.BatchNorm2d(planes * self.expansion)
self.relu = nn.ReLU(inplace=True)
self.downsample = downsample
self.stride = stride
def forward(self, x):
identity = x
out = self.relu(self.bn1(self.conv1(x)))
out = self.relu(self.bn2(self.conv2(out)))
out = self.bn3(self.conv3(out))
if self.downsample is not None:
identity = self.downsample(x)
out += identity
out = self.relu(out)
return out
class DeepfakeClassifier(nn.Module):
def __init__(self):
super(DeepfakeClassifier, self).__init__()
self.model = nn.Sequential(
nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1),
self._make_layer(ResNeXtBottleneck, 128, 3, stride=1), # 128*2 = 256
self._make_layer(ResNeXtBottleneck, 256, 4, stride=2), # 256*2 = 512
self._make_layer(ResNeXtBottleneck, 512, 6, stride=2), # 512*2 = 1024
self._make_layer(ResNeXtBottleneck, 1024, 3, stride=2) # 1024*2 = 2048
)
self.lstm = nn.LSTM(input_size=2048, hidden_size=2048, batch_first=True, bias=True)
self.linear1 = nn.Linear(2048, 2)
def _make_layer(self, block, planes, blocks, stride=1):
inplanes = self._get_inplanes()
downsample = None
if stride != 1 or inplanes != planes * block.expansion:
downsample = nn.Sequential(
nn.Conv2d(inplanes, planes * block.expansion, kernel_size=1, stride=stride, bias=False),
nn.BatchNorm2d(planes * block.expansion),
)
layers = []
layers.append(block(inplanes, planes, stride, downsample))
self._inplanes = planes * block.expansion
for _ in range(1, blocks):
layers.append(block(self._inplanes, planes))
return nn.Sequential(*layers)
def _get_inplanes(self):
if not hasattr(self, '_inplanes'):
self._inplanes = 64
return self._inplanes
def forward(self, x):
if x.ndim == 4:
# Single image input: (batch, channels, height, width)
x = self.model(x)
x = nn.functional.adaptive_avg_pool2d(x, (1, 1))
x = x.view(x.size(0), -1)
return self.linear1(x) # No LSTM for images
elif x.ndim == 5:
# Video input: (batch, frames, channels, height, width)
batch, frames, c, h, w = x.shape
x = x.view(-1, c, h, w)
x = self.model(x)
x = nn.functional.adaptive_avg_pool2d(x, (1, 1))
x = x.view(batch, frames, -1)
x, _ = self.lstm(x)
x = x[:, -1, :]
return self.linear1(x)
else:
raise ValueError(f"Unexpected input shape: {x.shape}")