(self,fea_dim)
| 50 | |
| 51 | class bVGG(torch.nn.Module): |
| 52 | def __init__(self,fea_dim): |
| 53 | super(bVGG, self).__init__() |
| 54 | # self.img_dim = 4096 |
| 55 | self.img_dim = 2048 |
| 56 | self.attention = Attention(dim=128,heads=4) |
| 57 | |
| 58 | self.linear_img = nn.Sequential(torch.nn.Linear(self.img_dim, fea_dim),torch.nn.ReLU()) |
| 59 | |
| 60 | self.classifier = nn.Linear(fea_dim,2) |
| 61 | |
| 62 | def forward(self, **kwargs): |
| 63 | frames=kwargs['frames'] |