forked from yihuacheng/Gaze-Net
-
Notifications
You must be signed in to change notification settings - Fork 0
/
Copy pathmodel.py
61 lines (48 loc) · 1.69 KB
/
model.py
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
import math
import numpy as np
class model(nn.Module):
def __init__(self):
super(model, self).__init__()
vgg16 = torchvision.models.vgg16(pretrained=True)
self.convNet = vgg16.features
self.FC = nn.Sequential(
nn.Linear(512*4*7, 4096),
nn.ReLU(inplace=True),
nn.Dropout(0.5)
)
self.output = nn.Sequential(
nn.Linear(4096+2, 4096),
nn.ReLU(inplace=True),
nn.Dropout(0.5),
nn.Linear(4096, 2),
)
# replace the maxpooling layer in VGG
self.convNet[4] = nn.MaxPool2d(kernel_size=2, stride=1)
self.convNet[9] = nn.MaxPool2d(kernel_size=2, stride=1)
def forward(self, x_in):
feature = self.convNet(x_in['eye'])
feature = torch.flatten(feature, start_dim=1)
feature = self.FC(feature)
feature = torch.cat((feature, x_in['head_pose']), 1)
gaze = self.output(feature)
return gaze
def _init_weights(self):
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.kaiming_uniform_(m.weight, mode="fan_in", nonlinearity="relu")
nn.init.zeros_(m.bias)
if __name__ == '__main__':
m = model().cuda()
'''feature = {"face":torch.zeros(10, 3, 224, 224).cuda(),
"left":torch.zeros(10,1, 36,60).cuda(),
"right":torch.zeros(10,1, 36,60).cuda()
}'''
feature = {"head_pose": torch.zeros(10, 2).cuda(),
"eye": torch.zeros(10, 3, 36, 60).cuda()
}
a = m(feature)
print(m)