Skip to content
This repository was archived by the owner on Jan 15, 2024. It is now read-only.

Commit ab24ce6

Browse files
zhresholdszha
authored andcommitted
remove _fork (#353)
* remove _fork * No rely on latest pr * remove _fork * No rely on latest pr * add unittest for record file * add comments * remove _fork * No rely on latest pr * add unittest for record file * add comments * mark test as serial * disable record * fix test case
1 parent f81042d commit ab24ce6

2 files changed

Lines changed: 46 additions & 2 deletions

File tree

gluonnlp/data/dataloader.py

Lines changed: 22 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,13 +19,33 @@
1919
"""DataLoader. An extension of Gluon data loader that allows multi-shard sampling."""
2020
__all__ = ['ShardedDataLoader']
2121

22+
import sys
2223
from mxnet import context
23-
from mxnet.gluon.data.dataloader import DataLoader, _MultiWorkerIter, _as_in_context
24+
from mxnet.gluon.data.dataloader import DataLoader
25+
from mxnet.gluon.data.dataloader import _MultiWorkerIter, _as_in_context
26+
from mxnet.recordio import MXRecordIO
27+
28+
def _recursive_fork_recordio(obj, depth, max_depth=1000):
29+
"""Recursively find instance of MXRecordIO and reset file handler.
30+
This is required for MXRecordIO which holds a C pointer to a opened file after fork.
31+
"""
32+
if depth >= max_depth:
33+
return
34+
if isinstance(obj, MXRecordIO):
35+
obj.close()
36+
obj.open() # re-obtain file hanlder in new process
37+
elif (hasattr(obj, '__dict__')):
38+
for _, v in obj.__dict__.items():
39+
_recursive_fork_recordio(v, depth + 1, max_depth)
2440

2541

2642
def worker_loop(dataset, key_queue, data_queue, batchify_fn):
2743
"""Worker loop for multiprocessing DataLoader."""
28-
dataset._fork()
44+
# re-fork a new recordio handler in new process if applicable
45+
limit = sys.getrecursionlimit()
46+
max_recursion_depth = min(limit - 5, max(10, limit // 2))
47+
_recursive_fork_recordio(dataset, 0, max_recursion_depth)
48+
2949
while True:
3050
idx, samples = key_queue.get()
3151
if idx is None:

tests/unittest/train/test_dataloader.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
import numpy as np
2+
import os
23
import mxnet as mx
34
from gluonnlp.data import FixedBucketSampler, ShardedDataLoader
45
from mxnet import gluon
6+
from mxnet.gluon.utils import download
57
import pytest
68

79

@@ -36,3 +38,25 @@ def test_sharded_data_loader():
3638
assert mx.test_utils.almost_equal(seqs[j][1].asnumpy(),
3739
Y[(i*num_shards+j)*2-num_shards:
3840
(i*num_shards+j+1)*2-num_shards])
41+
42+
def test_sharded_data_loader_record_file():
43+
# test record file
44+
url_format = 'https://apache-mxnet.s3-accelerate.amazonaws.com/gluon/dataset/pikachu/{}'
45+
filename = 'val.rec'
46+
idx_filename = 'val.idx'
47+
download(url_format.format(filename), path=os.path.join('tests', 'data', filename))
48+
download(url_format.format(idx_filename), path=os.path.join('tests', 'data', idx_filename))
49+
rec_dataset = gluon.data.vision.ImageRecordDataset(os.path.join('tests', 'data', filename))
50+
51+
num_workers = 2
52+
num_shards = 4
53+
X = np.random.uniform(size=(100, 20))
54+
Y = np.random.uniform(size=(100,))
55+
batch_sampler = FixedBucketSampler(lengths=[X.shape[1]] * X.shape[0],
56+
batch_size=2,
57+
num_buckets=1,
58+
shuffle=False,
59+
num_shards=num_shards)
60+
loader = ShardedDataLoader(rec_dataset, batch_sampler=batch_sampler, num_workers=num_workers)
61+
for i, seqs in enumerate(loader):
62+
assert len(seqs) == num_shards

0 commit comments

Comments
 (0)