MCPcopy Create free account
hub / github.com/ICTMCG/FakeSV / TikTecModel

Class TikTecModel

code/models/TikTec.py:114–140  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

112 return video_rep
113
114class 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

Callers 1

get_modelMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected