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

Method decode

openrec/modeling/decoders/parseq_decoder.py:244–265  ·  view source on GitHub ↗
(
        self,
        tgt: torch.Tensor,
        memory: torch.Tensor,
        tgt_mask: Optional[Tensor] = None,
        tgt_padding_mask: Optional[Tensor] = None,
        tgt_query: Optional[Tensor] = None,
        tgt_query_mask: Optional[Tensor] = None,
        pos_query: torch.Tensor = None,
    )

Source from the content-addressed store, hash-verified

242 return param_names
243
244 def decode(
245 self,
246 tgt: torch.Tensor,
247 memory: torch.Tensor,
248 tgt_mask: Optional[Tensor] = None,
249 tgt_padding_mask: Optional[Tensor] = None,
250 tgt_query: Optional[Tensor] = None,
251 tgt_query_mask: Optional[Tensor] = None,
252 pos_query: torch.Tensor = None,
253 ):
254 N, L = tgt.shape
255 # <bos> stands for the null context. We only supply position information for characters after <bos>.
256 null_ctx = self.text_embed(tgt[:, :1])
257
258 if tgt_query is None:
259 tgt_query = pos_query[:, :L]
260 tgt_emb = pos_query[:, :L - 1] + self.text_embed(tgt[:, 1:])
261 tgt_emb = self.dropout(torch.cat([null_ctx, tgt_emb], dim=1))
262
263 tgt_query = self.dropout(tgt_query)
264 return self.decoder(tgt_query, tgt_emb, memory, tgt_query_mask,
265 tgt_mask, tgt_padding_mask)
266
267 def forward(self, x, data=None, pos_query=None):
268 if self.training:

Callers 2

forward_testMethod · 0.95
training_stepMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected