Samanth/I3D_Sign_Language_Classification
0
1import torch2import torch.nn as nn3import torch.nn.functional as F4from torch.autograd import Variable5 6import numpy as np7 8import os9import sys10from collections import OrderedDict11 12 13class MaxPool3dSamePadding(nn.MaxPool3d):14 15 def compute_pad(self, dim, s):16 if s % self.stride[dim] == 0:17 return max(self.kernel_size[dim] - self.stride[dim], 0)18 else:19 return max(self.kernel_size[dim] - (s % self.stride[dim]), 0)20 21 def forward(self, x):22 # compute 'same' padding23 (batch, channel, t, h, w) = x.size()24 #print t,h,w25 out_t = np.ceil(float(t) / float(self.stride[0]))26 out_h = np.ceil(float(h) / float(self.stride[1]))27 out_w = np.ceil(float(w) / float(self.stride[2]))28 #print out_t, out_h, out_w29 pad_t = self.compute_pad(0, t)30 pad_h = self.compute_pad(1, h)31 pad_w = self.compute_pad(2, w)32 #print pad_t, pad_h, pad_w33 34 pad_t_f = pad_t // 235 pad_t_b = pad_t - pad_t_f36 pad_h_f = pad_h // 237 pad_h_b = pad_h - pad_h_f38 pad_w_f = pad_w // 239 pad_w_b = pad_w - pad_w_f40 41 pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)42 #print x.size()43 #print pad44 x = F.pad(x, pad)45 return super(MaxPool3dSamePadding, self).forward(x)46 47 48class Unit3D(nn.Module):49 50 def __init__(self, in_channels,51 output_channels,52 kernel_shape=(1, 1, 1),53 stride=(1, 1, 1),54 padding=0,55 activation_fn=F.relu,56 use_batch_norm=True,57 use_bias=False,58 name='unit_3d'):59 60 """Initializes Unit3D module."""61 super(Unit3D, self).__init__()62 63 self._output_channels = output_channels64 self._kernel_shape = kernel_shape65 self._stride = stride66 self._use_batch_norm = use_batch_norm67 self._activation_fn = activation_fn68 self._use_bias = use_bias69 self.name = name70 self.padding = padding71 72 self.conv3d = nn.Conv3d(in_channels=in_channels,73 out_channels=self._output_channels,74 kernel_size=self._kernel_shape,75 stride=self._stride,76 padding=0, # we always want padding to be 0 here. We will dynamically pad based on input size in forward function77 bias=self._use_bias)78 79 if self._use_batch_norm:80 self.bn = nn.BatchNorm3d(self._output_channels, eps=0.001, momentum=0.01)81 82 def compute_pad(self, dim, s):83 if s % self._stride[dim] == 0:84 return max(self._kernel_shape[dim] - self._stride[dim], 0)85 else:86 return max(self._kernel_shape[dim] - (s % self._stride[dim]), 0)87 88 89 def forward(self, x):90 # compute 'same' padding91 (batch, channel, t, h, w) = x.size()92 #print t,h,w93 out_t = np.ceil(float(t) / float(self._stride[0]))94 out_h = np.ceil(float(h) / float(self._stride[1]))95 out_w = np.ceil(float(w) / float(self._stride[2]))96 #print out_t, out_h, out_w97 pad_t = self.compute_pad(0, t)98 pad_h = self.compute_pad(1, h)99 pad_w = self.compute_pad(2, w)100 #print pad_t, pad_h, pad_w101 102 pad_t_f = pad_t // 2103 pad_t_b = pad_t - pad_t_f104 pad_h_f = pad_h // 2105 pad_h_b = pad_h - pad_h_f106 pad_w_f = pad_w // 2107 pad_w_b = pad_w - pad_w_f108 109 pad = (pad_w_f, pad_w_b, pad_h_f, pad_h_b, pad_t_f, pad_t_b)110 #print x.size()111 #print pad112 x = F.pad(x, pad)113 #print x.size() 114 115 x = self.conv3d(x)116 if self._use_batch_norm:117 x = self.bn(x)118 if self._activation_fn is not None:119 x = self._activation_fn(x)120 return x121 122 123 124class InceptionModule(nn.Module):125 def __init__(self, in_channels, out_channels, name):126 super(InceptionModule, self).__init__()127 128 self.b0 = Unit3D(in_channels=in_channels, output_channels=out_channels[0], kernel_shape=[1, 1, 1], padding=0,129 name=name+'/Branch_0/Conv3d_0a_1x1')130 self.b1a = Unit3D(in_channels=in_channels, output_channels=out_channels[1], kernel_shape=[1, 1, 1], padding=0,131 name=name+'/Branch_1/Conv3d_0a_1x1')132 self.b1b = Unit3D(in_channels=out_channels[1], output_channels=out_channels[2], kernel_shape=[3, 3, 3],133 name=name+'/Branch_1/Conv3d_0b_3x3')134 self.b2a = Unit3D(in_channels=in_channels, output_channels=out_channels[3], kernel_shape=[1, 1, 1], padding=0,135 name=name+'/Branch_2/Conv3d_0a_1x1')136 self.b2b = Unit3D(in_channels=out_channels[3], output_channels=out_channels[4], kernel_shape=[3, 3, 3],137 name=name+'/Branch_2/Conv3d_0b_3x3')138 self.b3a = MaxPool3dSamePadding(kernel_size=[3, 3, 3],139 stride=(1, 1, 1), padding=0)140 self.b3b = Unit3D(in_channels=in_channels, output_channels=out_channels[5], kernel_shape=[1, 1, 1], padding=0,141 name=name+'/Branch_3/Conv3d_0b_1x1')142 self.name = name143 144 def forward(self, x): 145 b0 = self.b0(x)146 b1 = self.b1b(self.b1a(x))147 b2 = self.b2b(self.b2a(x))148 b3 = self.b3b(self.b3a(x))149 return torch.cat([b0,b1,b2,b3], dim=1)150 151 152class InceptionI3d(nn.Module):153 """Inception-v1 I3D architecture.154 The model is introduced in:155 Quo Vadis, Action Recognition? A New Model and the Kinetics Dataset156 Joao Carreira, Andrew Zisserman157 https://arxiv.org/pdf/1705.07750v1.pdf.158 See also the Inception architecture, introduced in:159 Going deeper with convolutions160 Christian Szegedy, Wei Liu, Yangqing Jia, Pierre Sermanet, Scott Reed,161 Dragomir Anguelov, Dumitru Erhan, Vincent Vanhoucke, Andrew Rabinovich.162 http://arxiv.org/pdf/1409.4842v1.pdf.163 """164 165 # Endpoints of the model in order. During construction, all the endpoints up166 # to a designated `final_endpoint` are returned in a dictionary as the167 # second return value.168 VALID_ENDPOINTS = (169 'Conv3d_1a_7x7',170 'MaxPool3d_2a_3x3',171 'Conv3d_2b_1x1',172 'Conv3d_2c_3x3',173 'MaxPool3d_3a_3x3',174 'Mixed_3b',175 'Mixed_3c',176 'MaxPool3d_4a_3x3',177 'Mixed_4b',178 'Mixed_4c',179 'Mixed_4d',180 'Mixed_4e',181 'Mixed_4f',182 'MaxPool3d_5a_2x2',183 'Mixed_5b',184 'Mixed_5c',185 'Logits',186 'Predictions',187 )188 189 def __init__(self, num_classes=400, spatial_squeeze=True,190 final_endpoint='Logits', name='inception_i3d', in_channels=3, dropout_keep_prob=0.5):191 """Initializes I3D model instance.192 Args:193 num_classes: The number of outputs in the logit layer (default 400, which194 matches the Kinetics dataset).195 spatial_squeeze: Whether to squeeze the spatial dimensions for the logits196 before returning (default True).197 final_endpoint: The model contains many possible endpoints.198 `final_endpoint` specifies the last endpoint for the model to be built199 up to. In addition to the output at `final_endpoint`, all the outputs200 at endpoints up to `final_endpoint` will also be returned, in a201 dictionary. `final_endpoint` must be one of202 InceptionI3d.VALID_ENDPOINTS (default 'Logits').203 name: A string (optional). The name of this module.204 Raises:205 ValueError: if `final_endpoint` is not recognized.206 """207 208 if final_endpoint not in self.VALID_ENDPOINTS:209 raise ValueError('Unknown final endpoint %s' % final_endpoint)210 211 super(InceptionI3d, self).__init__()212 self._num_classes = num_classes213 self._spatial_squeeze = spatial_squeeze214 self._final_endpoint = final_endpoint215 self.logits = None216 217 if self._final_endpoint not in self.VALID_ENDPOINTS:218 raise ValueError('Unknown final endpoint %s' % self._final_endpoint)219 220 self.end_points = {}221 end_point = 'Conv3d_1a_7x7'222 self.end_points[end_point] = Unit3D(in_channels=in_channels, output_channels=64, kernel_shape=[7, 7, 7],223 stride=(2, 2, 2), padding=(3,3,3), name=name+end_point)224 if self._final_endpoint == end_point: return225 226 end_point = 'MaxPool3d_2a_3x3'227 self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),228 padding=0)229 if self._final_endpoint == end_point: return230 231 end_point = 'Conv3d_2b_1x1'232 self.end_points[end_point] = Unit3D(in_channels=64, output_channels=64, kernel_shape=[1, 1, 1], padding=0,233 name=name+end_point)234 if self._final_endpoint == end_point: return235 236 end_point = 'Conv3d_2c_3x3'237 self.end_points[end_point] = Unit3D(in_channels=64, output_channels=192, kernel_shape=[3, 3, 3], padding=1,238 name=name+end_point)239 if self._final_endpoint == end_point: return240 241 end_point = 'MaxPool3d_3a_3x3'242 self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[1, 3, 3], stride=(1, 2, 2),243 padding=0)244 if self._final_endpoint == end_point: return245 246 end_point = 'Mixed_3b'247 self.end_points[end_point] = InceptionModule(192, [64,96,128,16,32,32], name+end_point)248 if self._final_endpoint == end_point: return249 250 end_point = 'Mixed_3c'251 self.end_points[end_point] = InceptionModule(256, [128,128,192,32,96,64], name+end_point)252 if self._final_endpoint == end_point: return253 254 end_point = 'MaxPool3d_4a_3x3'255 self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[3, 3, 3], stride=(2, 2, 2),256 padding=0)257 if self._final_endpoint == end_point: return258 259 end_point = 'Mixed_4b'260 self.end_points[end_point] = InceptionModule(128+192+96+64, [192,96,208,16,48,64], name+end_point)261 if self._final_endpoint == end_point: return262 263 end_point = 'Mixed_4c'264 self.end_points[end_point] = InceptionModule(192+208+48+64, [160,112,224,24,64,64], name+end_point)265 if self._final_endpoint == end_point: return266 267 end_point = 'Mixed_4d'268 self.end_points[end_point] = InceptionModule(160+224+64+64, [128,128,256,24,64,64], name+end_point)269 if self._final_endpoint == end_point: return270 271 end_point = 'Mixed_4e'272 self.end_points[end_point] = InceptionModule(128+256+64+64, [112,144,288,32,64,64], name+end_point)273 if self._final_endpoint == end_point: return274 275 end_point = 'Mixed_4f'276 self.end_points[end_point] = InceptionModule(112+288+64+64, [256,160,320,32,128,128], name+end_point)277 if self._final_endpoint == end_point: return278 279 end_point = 'MaxPool3d_5a_2x2'280 self.end_points[end_point] = MaxPool3dSamePadding(kernel_size=[2, 2, 2], stride=(2, 2, 2),281 padding=0)282 if self._final_endpoint == end_point: return283 284 end_point = 'Mixed_5b'285 self.end_points[end_point] = InceptionModule(256+320+128+128, [256,160,320,32,128,128], name+end_point)286 if self._final_endpoint == end_point: return287 288 end_point = 'Mixed_5c'289 self.end_points[end_point] = InceptionModule(256+320+128+128, [384,192,384,48,128,128], name+end_point)290 if self._final_endpoint == end_point: return291 292 end_point = 'Logits'293 self.avg_pool = nn.AvgPool3d(kernel_size=[2, 7, 7],294 stride=(1, 1, 1))295 self.dropout = nn.Dropout(dropout_keep_prob)296 self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,297 kernel_shape=[1, 1, 1],298 padding=0,299 activation_fn=None,300 use_batch_norm=False,301 use_bias=True,302 name='logits')303 304 self.build()305 306 307 def replace_logits(self, num_classes):308 self._num_classes = num_classes309 self.logits = Unit3D(in_channels=384+384+128+128, output_channels=self._num_classes,310 kernel_shape=[1, 1, 1],311 padding=0,312 activation_fn=None,313 use_batch_norm=False,314 use_bias=True,315 name='logits')316 317 def build(self):318 for k in self.end_points.keys():319 self.add_module(k, self.end_points[k])320 321 def forward(self, x, pretrained=False, n_tune_layers=-1):322 if pretrained:323 assert n_tune_layers >= 0324 325 freeze_endpoints = self.VALID_ENDPOINTS[:-n_tune_layers]326 tune_endpoints = self.VALID_ENDPOINTS[-n_tune_layers:]327 else:328 freeze_endpoints = []329 tune_endpoints = self.VALID_ENDPOINTS330 331 # backbone, no gradient part332 with torch.no_grad():333 for end_point in freeze_endpoints:334 if end_point in self.end_points:335 x = self._modules[end_point](x) # use _modules to work with dataparallel336 337 # backbone, gradient part338 for end_point in tune_endpoints:339 if end_point in self.end_points:340 x = self._modules[end_point](x) # use _modules to work with dataparallel341 342 # head343 x = self.logits(self.dropout(self.avg_pool(x)))344 if self._spatial_squeeze:345 logits = x.squeeze(3).squeeze(3)346 # logits is batch X time X classes, which is what we want to work with347 return logits348 349 350 def extract_features(self, x):351 for end_point in self.VALID_ENDPOINTS:352 if end_point in self.end_points:353 x = self._modules[end_point](x)354 return self.avg_pool(x)