MCPcopy Create free account
hub / github.com/chibohe/text_recognition_toolbox / forward

Method forward

networks/SAR.py:183–208  ·  view source on GitHub ↗
(self, vis_features, holistic_features, text, num_steps=30, is_train=True)

Source from the content-addressed store, hash-verified

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
211class AttentionCell(nn.Module):

Callers

nothing calls this directly

Calls 1

_char_one_hotMethod · 0.95

Tested by

no test coverage detected