GSL-video2text / utils /decode.py
tinh2312's picture
init
157fb39
Raw
History Blame Contribute Delete
1.33 kB
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