Skip to content

Commit 736c473

Browse files
authored
Fix interleave support in TensorFlow 1.14.
1 parent 84dac3d commit 736c473

3 files changed

Lines changed: 59 additions & 11 deletions

File tree

docs/conf.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@
2424
project = 'HybridBackend'
2525
author = 'Alibaba Group Holding Limited'
2626
copyright = '2021 Alibaba Group Holding Limited' # pylint: disable=redefined-builtin
27-
release = '0.1.0'
27+
release = 'latest'
2828

2929
# -- General configuration ---------------------------------------------------
3030

hybridbackend/tensorflow/data/parquet_dataset_v1.py

Lines changed: 4 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,6 @@
2323
from __future__ import print_function
2424

2525
from tensorflow.python.data.ops import dataset_ops
26-
from tensorflow.python.data.ops import readers
2726
from tensorflow.python.framework import dtypes
2827
from tensorflow.python.framework import ops
2928
from tensorflow.python.util import nest
@@ -186,13 +185,8 @@ def _build_dataset(
186185
'''
187186
if num_parallel_reads is None:
188187
return filenames.flat_map(dataset_creator)
189-
if num_parallel_reads == dataset_ops.AUTOTUNE:
190-
return filenames.interleave(
191-
dataset_creator, num_parallel_calls=num_parallel_reads)
192-
return readers.ParallelInterleaveDataset(
193-
filenames, dataset_creator,
194-
cycle_length=num_parallel_reads,
188+
return filenames.interleave(
189+
dataset_creator,
190+
cycle_length=num_parallel_reads if num_parallel_reads > 0 else 1,
195191
block_length=num_sequential_reads,
196-
sloppy=True,
197-
buffer_output_elements=None,
198-
prefetch_input_elements=1)
192+
num_parallel_calls=num_parallel_reads)

tests/tensorflow/data/parquet_dataset_test.py

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,10 @@
2626
import tempfile
2727

2828
from tensorflow.python.data.ops.dataset_ops import Dataset
29+
try:
30+
from tensorflow.python.data.ops.dataset_ops import AUTOTUNE
31+
except ImportError:
32+
from tensorflow.python.data.experimental.ops.optimization import AUTOTUNE
2933
from tensorflow.python.framework import dtypes
3034
from tensorflow.python.framework import errors
3135
from tensorflow.python.framework import ops
@@ -129,6 +133,56 @@ def gen_filenames():
129133
with self.assertRaises(errors.OutOfRangeError):
130134
sess.run(batch)
131135

136+
def test_read_from_generator_parallel(self):
137+
num_epochs = 2
138+
batch_size = 100
139+
with ops.Graph().as_default() as graph:
140+
def gen_filenames():
141+
for i in range(num_epochs + 1):
142+
if i == num_epochs:
143+
return # raise StopIteration
144+
yield self._filename
145+
filenames = Dataset.from_generator(
146+
gen_filenames, dtypes.string, tensor_shape.TensorShape([]))
147+
fields = [
148+
DataFrame.Field('A', dtypes.int64),
149+
DataFrame.Field('C', dtypes.int64)]
150+
ds = filenames.apply(
151+
read_parquet(batch_size, fields=fields, num_parallel_reads=3))
152+
ds = ds.prefetch(4)
153+
batch = make_one_shot_iterator(ds).get_next()
154+
155+
with self.test_session(use_gpu=False, graph=graph) as sess:
156+
for _ in range(len(self._df) * num_epochs // batch_size):
157+
sess.run(batch)
158+
with self.assertRaises(errors.OutOfRangeError):
159+
sess.run(batch)
160+
161+
def test_read_from_generator_parallel_auto(self):
162+
num_epochs = 2
163+
batch_size = 100
164+
with ops.Graph().as_default() as graph:
165+
def gen_filenames():
166+
for i in range(num_epochs + 1):
167+
if i == num_epochs:
168+
return # raise StopIteration
169+
yield self._filename
170+
filenames = Dataset.from_generator(
171+
gen_filenames, dtypes.string, tensor_shape.TensorShape([]))
172+
fields = [
173+
DataFrame.Field('A', dtypes.int64),
174+
DataFrame.Field('C', dtypes.int64)]
175+
ds = filenames.apply(
176+
read_parquet(batch_size, fields=fields, num_parallel_reads=AUTOTUNE))
177+
ds = ds.prefetch(4)
178+
batch = make_one_shot_iterator(ds).get_next()
179+
180+
with self.test_session(use_gpu=False, graph=graph) as sess:
181+
for _ in range(len(self._df) * num_epochs // batch_size):
182+
sess.run(batch)
183+
with self.assertRaises(errors.OutOfRangeError):
184+
sess.run(batch)
185+
132186

133187
if __name__ == '__main__':
134188
os.environ['CUDA_VISIBLE_DEVICES'] = ''

0 commit comments

Comments
 (0)