tf.data.dataset from avro files

Viewed 717

I'm trying to parallelize my input pipeline using tf.data.Dataset and TFRecordDataset.

files = tf.data.Dataset.list_files("./data/*.avro")
dataset = tf.data.TFRecordDataset(files, num_parallel_reads=16)
dataset = dataset.apply(tf.contrib.data.map_and_batch(
    preprocess_fn, 512, num_parallel_batches=16) )

I'm not sure how to write preprocess_fn if the input is an AVRO file (which is like JSON).


Currently, I am using tf.data.Dataset.from_generator and feeding it avro records parsed by pyavroc or similar avro readers. But I'm not sure how to parallelize this as from_generator method does not have num_parallel_reads option available.

def gen():
    for file in all_avro_files:
        x, y = read_local_avro_data(file)
        for i, sample in enumerate( x ):
            yield sample, y[i]

dataset = tf.data.Dataset.from_generator( gen, 
            (tf.float32, tf.float64),
            ( tf.TensorShape([13000]), tf.TensorShape([]) 
        ) 
    )

Reading file by file is clearly a bottleneck and I see all cores waiting for data after exhausting the data from previous batch.

How to optimize either of the approaches?

0 Answers
Related