Skip to content

Commit a478d37

Browse files
zhresholdUbuntu
authored andcommitted
Fix lazy record io when used with dataloader and multi_worker > 0 (apache#12554)
* temp solution to record file dataset with multi worker * fix cascaded dataset for gluon dataloader, when multi_worker > 0 is used
1 parent 856b503 commit a478d37

3 files changed

Lines changed: 29 additions & 11 deletions

File tree

python/mxnet/gluon/data/dataloader.py

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636

3737
from . import sampler as _sampler
3838
from ... import nd, context
39+
from ...recordio import MXRecordIO
3940

4041
if sys.platform == 'darwin' or sys.platform == 'win32':
4142
def rebuild_ndarray(*args):
@@ -158,10 +159,24 @@ def _as_in_context(data, ctx):
158159
return [_as_in_context(d, ctx) for d in data]
159160
return data
160161

162+
def _recursive_fork_recordio(obj, depth, max_depth=1000):
163+
"""Recursively find instance of MXRecordIO and reset file handler.
164+
This is required for MXRecordIO which holds a C pointer to a opened file after fork.
165+
"""
166+
if depth >= max_depth:
167+
return
168+
if isinstance(obj, MXRecordIO):
169+
obj.close()
170+
obj.open() # re-obtain file hanlder in new process
171+
elif (hasattr(obj, '__dict__')):
172+
for _, v in obj.__dict__.items():
173+
_recursive_fork_recordio(v, depth + 1, max_depth)
174+
161175
def worker_loop(dataset, key_queue, data_queue, batchify_fn):
162176
"""Worker loop for multiprocessing DataLoader."""
163-
if hasattr(dataset, '_fork') and callable(dataset._fork):
164-
dataset._fork()
177+
# re-fork a new recordio handler in new process if applicable
178+
_recursive_fork_recordio(dataset, 0, 1000)
179+
165180
while True:
166181
idx, samples = key_queue.get()
167182
if idx is None:
@@ -181,6 +196,7 @@ def fetcher_loop(data_queue, data_buffer, pin_memory=False):
181196
batch = _as_in_context(batch, context.cpu())
182197
data_buffer[idx] = batch
183198

199+
184200
class _MultiWorkerIter(object):
185201
"""Interal multi-worker iterator for DataLoader."""
186202
def __init__(self, num_workers, dataset, batchify_fn, batch_sampler, pin_memory=False,

python/mxnet/gluon/data/dataset.py

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -94,11 +94,6 @@ def base_fn(x, *args):
9494
return fn(x)
9595
return self.transform(base_fn, lazy)
9696

97-
def _fork(self):
98-
"""Protective operations required when launching multiprocess workers."""
99-
# for non file descriptor related datasets, just skip
100-
pass
101-
10297

10398
class SimpleDataset(Dataset):
10499
"""Simple Dataset wrapper for lists and arrays.
@@ -180,9 +175,6 @@ class RecordFileDataset(Dataset):
180175
def __init__(self, filename):
181176
self.idx_file = os.path.splitext(filename)[0] + '.idx'
182177
self.filename = filename
183-
self._fork()
184-
185-
def _fork(self):
186178
self._record = recordio.MXIndexedRecordIO(self.idx_file, self.filename, 'r')
187179

188180
def __getitem__(self, idx):

tests/python/unittest/test_gluon_data.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,8 @@ def prepare_record():
6565
@with_seed()
6666
def test_recordimage_dataset():
6767
recfile = prepare_record()
68-
dataset = gluon.data.vision.ImageRecordDataset(recfile)
68+
fn = lambda x, y : (x, y)
69+
dataset = gluon.data.vision.ImageRecordDataset(recfile).transform(fn)
6970
loader = gluon.data.DataLoader(dataset, 1)
7071

7172
for i, (x, y) in enumerate(loader):
@@ -84,6 +85,15 @@ def test_recordimage_dataset_with_data_loader_multiworker():
8485
assert x.shape[0] == 1 and x.shape[3] == 3
8586
assert y.asscalar() == i
8687

88+
# with transform
89+
fn = lambda x, y : (x, y)
90+
dataset = gluon.data.vision.ImageRecordDataset(recfile).transform(fn)
91+
loader = gluon.data.DataLoader(dataset, 1, num_workers=5)
92+
93+
for i, (x, y) in enumerate(loader):
94+
assert x.shape[0] == 1 and x.shape[3] == 3
95+
assert y.asscalar() == i
96+
8797
@with_seed()
8898
def test_sampler():
8999
seq_sampler = gluon.data.SequentialSampler(10)

0 commit comments

Comments
 (0)