TFRecord decode_raw for sequence feature

Viewed 180

I have such dataset in TFRecord format:

def _bytes_feature(value):
    if isinstance(value, type(tf.constant(0))):
        value = value.numpy()

    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))

sequence_dict = {
    'frames': tf.train.FeatureList(feature=frames),
    "label": tf.train.FeatureList(feature=[tf.train.Feature(int64_list=tf.train.Int64List(value=[token])) for token in tokens]),
}
context_dict = {
    "frames_count": tf.train.Feature(int64_list=tf.train.Int64List(value=[frames_count])),
    "num_tokens": num_tokens,
}

sequence_context = tf.train.Features(feature=context_dict)
sequence_list = tf.train.FeatureLists(feature_list=sequence_dict)
example = tf.train.SequenceExample(context=sequence_context, feature_lists=sequence_list)
  • frames is a sequence of 112x112 gray images, represented by a list of results of _bytes_feature function.
  • label is a sequence of tokens.

My task is seq2seq, so the whole sequence of tokens corresponds to the whole sequence of frames (len(frames) != len(label)). The task is lip-reading, if that makes more sense.

I'm loading the dataset this way:

sequence_features = {
    'frames': tf.io.FixedLenSequenceFeature([], dtype=tf.string),
    "label": tf.io.FixedLenSequenceFeature([], dtype=tf.int64),
}
context_features = {
    "frames_count": tf.io.FixedLenFeature([], dtype=tf.int64),
    "num_tokens":  tf.io.FixedLenFeature([], dtype=tf.int64),
}
dataset = tf.data.TFRecordDataset("train-0.tfrecord")
dataset = dataset.map(_parse_function)
dataset = dataset.padded_batch(3)

The problem is I can't write _parse_function properly, so I can iterate over padded batches of sequences of tf.int8 tensors representing frames for one video along with batches of corresponding labels. Also I want to avoid using VarLenFeature, because sparse tensors don't play well with CTC loss on GPU.

Here is what I've tried:

def _parse_function(example_proto):
    context, sequence, _ = tf.io.parse_sequence_example(example_proto, context_features=context_features, sequence_features=sequence_features)
    image = tf.io.decode_raw(sequence["frames"], tf.int8)
    label = sequence["label"]

    return image, label

throws InvalidArgumentError: DecodeRaw requires input strings to all be the same size, but element 1 has size 2444 != 2456 [[{{node DecodeRaw}}]]

Changing parse_sequence_example to parse_single_sequence_example does not help and throws the same error

So the question is how should I modify _parse_function to make it return batches of tf.int8 frame sequences, shaped BxTxWxH, where B is a batch size and T is a sequence length?

0 Answers
Related