| 112 | return video_rep |
| 113 | |
| 114 | class TikTecModel(nn.Module): |
| 115 | def __init__(self, word_dim=300, mfcc_dim=650, visual_dim=1000, obj_num=45, CVRL_gru_dim=200, ASRL_gru_dim=500, VCIF_d_H=200, VCIF_gru_f_dim=200, VCIF_gru_w_dim=100, VCIF_dropout=0.2, MLP_hidden_dims=[512], MLP_dropout=0.2): |
| 116 | super(TikTecModel, self).__init__() |
| 117 | self.CVRL = CVRL(d_w=word_dim, d_f=visual_dim, obj_num=obj_num, gru_dim=CVRL_gru_dim) |
| 118 | self.ASRL = ASRL(d_w=(word_dim + mfcc_dim), gru_dim=ASRL_gru_dim) |
| 119 | self.VCIF = VCIF(d_f=visual_dim, d_w=2*ASRL_gru_dim, d_H=VCIF_d_H, gru_f_dim=VCIF_gru_f_dim, gru_w_dim=VCIF_gru_w_dim, dropout=VCIF_dropout) |
| 120 | self.MLP = MLP(VCIF_gru_f_dim + VCIF_gru_w_dim, MLP_hidden_dims, 2, MLP_dropout) |
| 121 | |
| 122 | def forward(self, **kwargs): |
| 123 | # IN: |
| 124 | # caption_feature: (bs, K, S, word_dim) = (bs, 200, 100, 300) |
| 125 | # visual_feature: (bs, K, obj_num, visual_dim) = (bs, 200, 45, 1000) |
| 126 | # asr_feature: (bs, N, word_dim + mfcc_dim) = (bs, 500, 300 + 650) |
| 127 | # mask_K: (bs, K) = (bs, 200) |
| 128 | # mask_N: (bs, N) = (bs, 500) |
| 129 | # OUT: (bs, 2) |
| 130 | caption_feature = kwargs['caption_feature'] |
| 131 | visual_feature = kwargs['visual_feature'] |
| 132 | asr_feature = kwargs['asr_feature'] |
| 133 | mask_K = kwargs['mask_K'] |
| 134 | mask_N = kwargs['mask_N'] |
| 135 | |
| 136 | frame_visual_rep = self.CVRL(caption_feature, visual_feature) |
| 137 | text_audio_rep = self.ASRL(asr_feature) |
| 138 | video_rep = self.VCIF(frame_visual_rep, text_audio_rep, mask_K, mask_N) |
| 139 | output = self.MLP(video_rep) |
| 140 | return output |