(self, vis_features, holistic_features, text, num_steps=30, is_train=True)
| 181 | return one_hot |
| 182 | |
| 183 | def forward(self, vis_features, holistic_features, text, num_steps=30, is_train=True): |
| 184 | batch_size = vis_features.size(0) |
| 185 | hidden = holistic_features |
| 186 | if is_train: |
| 187 | num_steps = text.size(1) |
| 188 | output_hiddens = torch.FloatTensor(batch_size, num_steps, \ |
| 189 | self.input_size+self.de_hidden_size).zero_().to(device) |
| 190 | for i in range(num_steps): |
| 191 | target = self._char_one_hot(text[:, i], self.num_classes) |
| 192 | hidden = self.rnn(target, hidden) |
| 193 | g = self.attn_cell(hidden[0], vis_features) |
| 194 | output_hiddens[:, i, :] = torch.cat([hidden[0], g], dim=1) |
| 195 | probs = self.generator(output_hiddens) |
| 196 | else: |
| 197 | probs = torch.FloatTensor(batch_size, num_steps, \ |
| 198 | self.num_classes).zero_().to(device) |
| 199 | target = torch.FloatTensor(batch_size, self.num_classes).zero_().to(device) |
| 200 | for i in range(num_steps): |
| 201 | hidden = self.rnn(target, hidden) |
| 202 | g = self.attn_cell(hidden[0], vis_features) |
| 203 | concat_feature = torch.cat([hidden[0], g], dim=1) |
| 204 | prob = self.generator(concat_feature) |
| 205 | probs[:, i, :] = prob |
| 206 | _, next_input = prob.max(axis=1) |
| 207 | target = self._char_one_hot(next_input, self.num_classes) |
| 208 | return probs |
| 209 | |
| 210 | |
| 211 | class AttentionCell(nn.Module): |
nothing calls this directly
no test coverage detected