| 304 | |
| 305 | |
| 306 | class DeciWatchTransformer(nn.Module): |
| 307 | |
| 308 | def __init__(self, |
| 309 | input_nc, |
| 310 | encoder_hidden_dim=512, |
| 311 | decoder_hidden_dim=512, |
| 312 | nhead=8, |
| 313 | num_encoder_layers=6, |
| 314 | num_decoder_layers=6, |
| 315 | dim_feedforward=2048, |
| 316 | dropout=0.1, |
| 317 | activation='relu', |
| 318 | pre_norm=False): |
| 319 | super(DeciWatchTransformer, self).__init__() |
| 320 | |
| 321 | self.joints_dim = input_nc |
| 322 | # bring in semantic (5 frames) temporal information into tokens |
| 323 | self.decoder_embed = nn.Conv1d( |
| 324 | self.joints_dim, |
| 325 | decoder_hidden_dim, |
| 326 | kernel_size=5, |
| 327 | stride=1, |
| 328 | padding=2) |
| 329 | |
| 330 | self.encoder_embed = nn.Linear(self.joints_dim, encoder_hidden_dim) |
| 331 | |
| 332 | encoder_layer = DeciWatchTransformerEncoderLayer( |
| 333 | encoder_hidden_dim, nhead, dim_feedforward, dropout, activation, |
| 334 | pre_norm) |
| 335 | encoder_norm = nn.LayerNorm(encoder_hidden_dim) if pre_norm else None |
| 336 | self.encoder = DeciWatchTransformerEncoder(encoder_layer, |
| 337 | num_encoder_layers, |
| 338 | encoder_norm) |
| 339 | |
| 340 | decoder_layer = DeciWatchTransformerDecoderLayer( |
| 341 | decoder_hidden_dim, nhead, dim_feedforward, dropout, activation, |
| 342 | pre_norm) |
| 343 | decoder_norm = nn.LayerNorm(decoder_hidden_dim) |
| 344 | self.decoder = DeciWatchTransformerDecoder(decoder_layer, |
| 345 | num_decoder_layers, |
| 346 | decoder_norm) |
| 347 | |
| 348 | self.decoder_joints_embed = nn.Linear(decoder_hidden_dim, |
| 349 | self.joints_dim) |
| 350 | self.encoder_joints_embed = nn.Linear(encoder_hidden_dim, |
| 351 | self.joints_dim) |
| 352 | |
| 353 | # reset parameters |
| 354 | for p in self.parameters(): |
| 355 | if p.dim() > 1: |
| 356 | nn.init.xavier_uniform_(p) |
| 357 | |
| 358 | self.encoder_hidden_dim = encoder_hidden_dim |
| 359 | self.decoder_hidden_dim = decoder_hidden_dim |
| 360 | |
| 361 | self.nhead = nhead |
| 362 | |
| 363 | def _generate_square_subsequent_mask(self, sz): |