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