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)
framesis a sequence of112x112gray images, represented by a list of results of_bytes_featurefunction.labelis 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?