File size: 1,334 Bytes
157fb39
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
import tensorflow as tf
from itertools import groupby
import numpy as np


def ctc_decode_func(tf_gloss_logits, input_lengths, beam_size):
    ctc_decode, _ = tf.nn.ctc_beam_search_decoder(
        inputs=tf_gloss_logits,
        sequence_length=input_lengths,
        beam_width=beam_size,
        top_paths=1,
    )
    ctc_decode = ctc_decode[0]
    tmp_gloss_sequences = [[] for _ in range(input_lengths.shape[0])]
    for (value_idx, dense_idx) in enumerate(ctc_decode.indices):
        tmp_gloss_sequences[dense_idx[0]].append(
            ctc_decode.values[value_idx].numpy() + 1
        )
    decoded_gloss_sequences = []
    for seq_idx in range(0, len(tmp_gloss_sequences)):
        decoded_gloss_sequences.append(
            [x[0] for x in groupby(tmp_gloss_sequences[seq_idx])]
        )
    return decoded_gloss_sequences


def decode(gloss_logits, beam_size, input_lengths):
    gloss_logits = gloss_logits.permute(1, 0, 2)  # T,B,V  [10,1,1124]
    gloss_logits = gloss_logits.cpu().detach().numpy()
    tf_gloss_logits = np.concatenate(
        (gloss_logits[:, :, 1:], gloss_logits[:, :, 0, None]),
        axis=-1,
    )
    decoded_gloss_sequences = ctc_decode_func(
        tf_gloss_logits=tf_gloss_logits,
        input_lengths=input_lengths,
        beam_size=beam_size
    )
    return decoded_gloss_sequences