MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / forward_test

Method forward_test

openrec/modeling/decoders/cppd_decoder.py:286–331  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

284 return self.forward_test(x)
285
286 def forward_test(self, x):
287 if not self.ds:
288 visual_feats = x + self.vis_pos_embed
289 else:
290 visual_feats = x
291 bs = visual_feats.shape[0]
292
293 pos_node_embed = self.pos_node_embed(
294 torch.arange(self.max_len).cuda(
295 x.get_device())).unsqueeze(0) + self.char_pos_embed
296 pos_node_embed = torch.tile(pos_node_embed, [bs, 1, 1])
297
298 char_vis_node_query = visual_feats
299 pos_vis_node_query = torch.concat([pos_node_embed, visual_feats], 1)
300
301 for char_decoder_layer, pos_decoder_layer in zip(
302 self.char_node_decoder, self.pos_node_decoder):
303 char_vis_node_query = char_decoder_layer(char_vis_node_query,
304 char_vis_node_query)
305 pos_vis_node_query = pos_decoder_layer(
306 pos_vis_node_query, pos_vis_node_query[:, self.max_len:, :])
307
308 pos_node_query = pos_vis_node_query[:, :self.max_len, :]
309
310 char_vis_feats = char_vis_node_query
311 # pos_vis_feats = pos_vis_node_query[:, self.max_len :, :]
312
313 # pos_node_feats = self.edge_decoder(
314 # pos_node_query, char_vis_feats, pos_vis_feats
315 # ) # B, 26, dim
316
317 pos_node_feats = pos_node_query
318 for layer_i in range(self.rec_layer_num):
319 rec_layer = self.edge_decoder[layer_i]
320 if (self.rec_layer_num + layer_i) % 2 == 0:
321 pos_node_feats = rec_layer(pos_node_feats, pos_node_feats,
322 self.self_mask)
323 else:
324 pos_node_feats = rec_layer(pos_node_feats, char_vis_feats)
325 edge_feats = self.edge_fc(pos_node_feats) # B, 26, 37
326
327 edge_logits = F.softmax(
328 edge_feats,
329 -1) # * F.sigmoid(pos_node_feats1.unsqueeze(-1)) # B, 26, 37
330
331 return edge_logits
332
333 def forward_train(self, x, targets=None):
334 if not self.ds:

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected