Skip to content
Projects
Groups
Snippets
Help
Loading...
Help
Support
Keyboard shortcuts
?
Submit feedback
Contribute to GitLab
Sign in
Toggle navigation
S
seminar-breakout
Project overview
Project overview
Details
Activity
Releases
Repository
Repository
Files
Commits
Branches
Tags
Contributors
Graph
Compare
Issues
0
Issues
0
List
Boards
Labels
Milestones
Merge Requests
0
Merge Requests
0
CI / CD
CI / CD
Pipelines
Jobs
Schedules
Analytics
Analytics
CI / CD
Repository
Value Stream
Wiki
Wiki
Members
Members
Collapse sidebar
Close sidebar
Activity
Graph
Create a new issue
Jobs
Commits
Issue Boards
Open sidebar
Shashank Suhas
seminar-breakout
Commits
50ff9036
You need to sign in or sign up before continuing.
Commit
50ff9036
authored
Aug 21, 2019
by
Yuxin Wu
Browse files
Options
Browse Files
Download
Email Patches
Plain Diff
Add MultiProcessMapAndBatchData
parent
ba9d1793
Changes
4
Hide whitespace changes
Inline
Side-by-side
Showing
4 changed files
with
187 additions
and
36 deletions
+187
-36
tensorpack/dataflow/common.py
tensorpack/dataflow/common.py
+25
-5
tensorpack/dataflow/parallel_map.py
tensorpack/dataflow/parallel_map.py
+123
-7
tensorpack/graph_builder/model_desc.py
tensorpack/graph_builder/model_desc.py
+9
-5
tensorpack/input_source/input_source.py
tensorpack/input_source/input_source.py
+30
-19
No files found.
tensorpack/dataflow/common.py
View file @
50ff9036
...
@@ -35,6 +35,11 @@ class TestDataSpeed(ProxyDataFlow):
...
@@ -35,6 +35,11 @@ class TestDataSpeed(ProxyDataFlow):
super
(
TestDataSpeed
,
self
)
.
__init__
(
ds
)
super
(
TestDataSpeed
,
self
)
.
__init__
(
ds
)
self
.
test_size
=
int
(
size
)
self
.
test_size
=
int
(
size
)
self
.
warmup
=
int
(
warmup
)
self
.
warmup
=
int
(
warmup
)
self
.
_reset_called
=
False
def
reset_state
(
self
):
self
.
_reset_called
=
True
super
(
TestDataSpeed
,
self
)
.
reset_state
()
def
__iter__
(
self
):
def
__iter__
(
self
):
""" Will run testing at the beginning, then produce data normally. """
""" Will run testing at the beginning, then produce data normally. """
...
@@ -46,7 +51,8 @@ class TestDataSpeed(ProxyDataFlow):
...
@@ -46,7 +51,8 @@ class TestDataSpeed(ProxyDataFlow):
"""
"""
Start testing with a progress bar.
Start testing with a progress bar.
"""
"""
self
.
ds
.
reset_state
()
if
not
self
.
_reset_called
:
self
.
ds
.
reset_state
()
itr
=
self
.
ds
.
__iter__
()
itr
=
self
.
ds
.
__iter__
()
if
self
.
warmup
:
if
self
.
warmup
:
for
_
in
tqdm
.
trange
(
self
.
warmup
,
**
get_tqdm_kwargs
()):
for
_
in
tqdm
.
trange
(
self
.
warmup
,
**
get_tqdm_kwargs
()):
...
@@ -91,6 +97,7 @@ class BatchData(ProxyDataFlow):
...
@@ -91,6 +97,7 @@ class BatchData(ProxyDataFlow):
except
NotImplementedError
:
except
NotImplementedError
:
pass
pass
self
.
batch_size
=
int
(
batch_size
)
self
.
batch_size
=
int
(
batch_size
)
assert
self
.
batch_size
>
0
self
.
remainder
=
remainder
self
.
remainder
=
remainder
self
.
use_list
=
use_list
self
.
use_list
=
use_list
...
@@ -111,10 +118,10 @@ class BatchData(ProxyDataFlow):
...
@@ -111,10 +118,10 @@ class BatchData(ProxyDataFlow):
for
data
in
self
.
ds
:
for
data
in
self
.
ds
:
holder
.
append
(
data
)
holder
.
append
(
data
)
if
len
(
holder
)
==
self
.
batch_size
:
if
len
(
holder
)
==
self
.
batch_size
:
yield
BatchData
.
_
aggregate_batch
(
holder
,
self
.
use_list
)
yield
BatchData
.
aggregate_batch
(
holder
,
self
.
use_list
)
del
holder
[:]
del
holder
[:]
if
self
.
remainder
and
len
(
holder
)
>
0
:
if
self
.
remainder
and
len
(
holder
)
>
0
:
yield
BatchData
.
_
aggregate_batch
(
holder
,
self
.
use_list
)
yield
BatchData
.
aggregate_batch
(
holder
,
self
.
use_list
)
@
staticmethod
@
staticmethod
def
_batch_numpy
(
data_list
):
def
_batch_numpy
(
data_list
):
...
@@ -146,7 +153,18 @@ class BatchData(ProxyDataFlow):
...
@@ -146,7 +153,18 @@ class BatchData(ProxyDataFlow):
pass
pass
@
staticmethod
@
staticmethod
def
_aggregate_batch
(
data_holder
,
use_list
=
False
):
def
aggregate_batch
(
data_holder
,
use_list
=
False
):
"""
Aggregate a list of datapoints to one batched datapoint.
Args:
data_holder (list[dp]): each dp is either a list or a dict.
use_list (bool): whether to batch data into a list or a numpy array.
Returns:
dp: either a list or a dict, depend on the inputs.
Each item is a batched version of the corresponding inputs.
"""
first_dp
=
data_holder
[
0
]
first_dp
=
data_holder
[
0
]
if
isinstance
(
first_dp
,
(
list
,
tuple
)):
if
isinstance
(
first_dp
,
(
list
,
tuple
)):
result
=
[]
result
=
[]
...
@@ -164,6 +182,8 @@ class BatchData(ProxyDataFlow):
...
@@ -164,6 +182,8 @@ class BatchData(ProxyDataFlow):
result
[
key
]
=
data_list
result
[
key
]
=
data_list
else
:
else
:
result
[
key
]
=
BatchData
.
_batch_numpy
(
data_list
)
result
[
key
]
=
BatchData
.
_batch_numpy
(
data_list
)
else
:
raise
ValueError
(
"Data point has to be list/tuple/dict. Got {}"
.
format
(
type
(
first_dp
)))
return
result
return
result
...
@@ -202,7 +222,7 @@ class BatchDataByShape(BatchData):
...
@@ -202,7 +222,7 @@ class BatchDataByShape(BatchData):
holder
=
self
.
holder
[
shp
]
holder
=
self
.
holder
[
shp
]
holder
.
append
(
dp
)
holder
.
append
(
dp
)
if
len
(
holder
)
==
self
.
batch_size
:
if
len
(
holder
)
==
self
.
batch_size
:
yield
BatchData
.
_
aggregate_batch
(
holder
)
yield
BatchData
.
aggregate_batch
(
holder
)
del
holder
[:]
del
holder
[:]
...
...
tensorpack/dataflow/parallel_map.py
View file @
50ff9036
...
@@ -12,11 +12,12 @@ from ..utils.concurrency import StoppableThread, enable_death_signal
...
@@ -12,11 +12,12 @@ from ..utils.concurrency import StoppableThread, enable_death_signal
from
..utils.serialize
import
dumps
,
loads
from
..utils.serialize
import
dumps
,
loads
from
..utils.develop
import
log_deprecated
from
..utils.develop
import
log_deprecated
from
.base
import
DataFlow
,
DataFlowReentrantGuard
,
ProxyDataFlow
from
.base
import
DataFlow
,
DataFlowReentrantGuard
,
ProxyDataFlow
from
.common
import
RepeatedData
from
.common
import
RepeatedData
,
BatchData
from
.parallel
import
_bind_guard
,
_get_pipe_name
,
_MultiProcessZMQDataFlow
,
_repeat_iter
,
_zmq_catch_error
from
.parallel
import
_bind_guard
,
_get_pipe_name
,
_MultiProcessZMQDataFlow
,
_repeat_iter
,
_zmq_catch_error
__all__
=
[
'MultiThreadMapData'
,
__all__
=
[
'MultiThreadMapData'
,
'MultiProcessMapData'
,
'MultiProcessMapDataZMQ'
]
'MultiProcessMapData'
,
'MultiProcessMapDataZMQ'
,
'MultiProcessMapAndBatchData'
,
'MultiProcessMapAndBatchDataZMQ'
]
class
_ParallelMapData
(
ProxyDataFlow
):
class
_ParallelMapData
(
ProxyDataFlow
):
...
@@ -286,6 +287,9 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
...
@@ -286,6 +287,9 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
self
.
_strict
=
strict
self
.
_strict
=
strict
self
.
_procs
=
[]
self
.
_procs
=
[]
def
_create_worker
(
self
,
id
,
pipename
,
hwm
):
return
MultiProcessMapDataZMQ
.
_Worker
(
id
,
self
.
map_func
,
pipename
,
hwm
)
def
reset_state
(
self
):
def
reset_state
(
self
):
_MultiProcessZMQDataFlow
.
reset_state
(
self
)
_MultiProcessZMQDataFlow
.
reset_state
(
self
)
_ParallelMapData
.
reset_state
(
self
)
_ParallelMapData
.
reset_state
(
self
)
...
@@ -299,9 +303,8 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
...
@@ -299,9 +303,8 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
self
.
_proc_ids
=
[
u'{}'
.
format
(
k
)
.
encode
(
'utf-8'
)
for
k
in
range
(
self
.
num_proc
)]
self
.
_proc_ids
=
[
u'{}'
.
format
(
k
)
.
encode
(
'utf-8'
)
for
k
in
range
(
self
.
num_proc
)]
worker_hwm
=
int
(
self
.
_buffer_size
*
2
//
self
.
num_proc
)
worker_hwm
=
int
(
self
.
_buffer_size
*
2
//
self
.
num_proc
)
self
.
_procs
=
[
MultiProcessMapDataZMQ
.
_Worker
(
self
.
_procs
=
[
self
.
_create_worker
(
self
.
_proc_ids
[
k
],
pipename
,
worker_hwm
)
self
.
_proc_ids
[
k
],
self
.
map_func
,
pipename
,
worker_hwm
)
for
k
in
range
(
self
.
num_proc
)]
for
k
in
range
(
self
.
num_proc
)]
self
.
_start_processes
()
self
.
_start_processes
()
self
.
_fill_buffer
()
# pre-fill the bufer
self
.
_fill_buffer
()
# pre-fill the bufer
...
@@ -316,12 +319,120 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
...
@@ -316,12 +319,120 @@ class MultiProcessMapDataZMQ(_ParallelMapData, _MultiProcessZMQDataFlow):
return
dp
return
dp
def
__iter__
(
self
):
def
__iter__
(
self
):
with
self
.
_guard
,
_zmq_catch_error
(
'MultiProcessMapData'
):
with
self
.
_guard
,
_zmq_catch_error
(
type
(
self
)
.
__name__
):
for
dp
in
super
(
MultiProcessMapDataZMQ
,
self
)
.
__iter__
():
for
dp
in
super
(
MultiProcessMapDataZMQ
,
self
)
.
__iter__
():
yield
dp
yield
dp
MultiProcessMapData
=
MultiProcessMapDataZMQ
# alias
class
MultiProcessMapAndBatchDataZMQ
(
_MultiProcessZMQDataFlow
):
"""
Similar to :class:`MultiProcessMapDataZMQ`, except that this DataFlow
also does batching in parallel in the worker processes.
Therefore it can be helpful if you wish to hide the latency of batching.
When `nr_proc==1`, the behavior of this class is identical to
`BatchData(MapData(ds, map_func), batch_size)`.
When `nr_proc>1`, the datapoints may be grouped in arbitrary order,
or grouped with datapoints from a different pass of the given dataflow.
"""
class
_Dispatcher
(
mp
.
Process
):
def
__init__
(
self
,
ds
,
pipename
,
hwm
):
super
(
MultiProcessMapAndBatchDataZMQ
.
_Dispatcher
,
self
)
.
__init__
()
self
.
ds
=
RepeatedData
(
ds
,
-
1
)
self
.
pipename
=
pipename
self
.
hwm
=
hwm
def
run
(
self
):
enable_death_signal
()
ctx
=
zmq
.
Context
()
socket
=
ctx
.
socket
(
zmq
.
PUSH
)
socket
.
set_hwm
(
self
.
hwm
)
socket
.
bind
(
self
.
pipename
)
self
.
ds
.
reset_state
()
for
dp
in
self
.
ds
:
socket
.
send
(
dumps
(
dp
),
copy
=
False
)
class
_Worker
(
mp
.
Process
):
def
__init__
(
self
,
identity
,
map_func
,
input_pipe
,
result_pipe
,
hwm
,
batch_size
):
super
(
MultiProcessMapAndBatchDataZMQ
.
_Worker
,
self
)
.
__init__
()
self
.
identity
=
identity
self
.
map_func
=
map_func
self
.
input_pipe
=
input_pipe
self
.
result_pipe
=
result_pipe
self
.
hwm
=
hwm
self
.
batch_size
=
batch_size
def
run
(
self
):
enable_death_signal
(
_warn
=
self
.
identity
==
b
'0'
)
ctx
=
zmq
.
Context
()
socket
=
ctx
.
socket
(
zmq
.
PULL
)
socket
.
setsockopt
(
zmq
.
IDENTITY
,
self
.
identity
)
socket
.
set_hwm
(
self
.
hwm
)
socket
.
connect
(
self
.
input_pipe
)
out_socket
=
ctx
.
socket
(
zmq
.
PUSH
)
out_socket
.
set_hwm
(
max
(
self
.
hwm
//
self
.
batch_size
,
5
))
out_socket
.
connect
(
self
.
result_pipe
)
batch
=
[]
while
True
:
dp
=
loads
(
socket
.
recv
(
copy
=
False
))
dp
=
self
.
map_func
(
dp
)
if
dp
is
not
None
:
batch
.
append
(
dp
)
if
len
(
batch
)
==
self
.
batch_size
:
dp
=
BatchData
.
aggregate_batch
(
batch
)
out_socket
.
send
(
dumps
(
dp
),
copy
=
False
)
del
batch
[:]
def
__init__
(
self
,
ds
,
num_proc
,
map_func
,
batch_size
,
buffer_size
=
1024
):
"""
Args:
ds (DataFlow): the dataflow to map
num_proc(int): number of threads to use
map_func (callable): datapoint -> datapoint | None. Return None to
discard/skip the datapoint.
batch_size (int): batch size
buffer_size (int): number of datapoints in the buffer
"""
super
(
MultiProcessMapAndBatchDataZMQ
,
self
)
.
__init__
()
self
.
ds
=
ds
self
.
num_proc
=
num_proc
self
.
map_func
=
map_func
self
.
buffer_size
=
buffer_size
self
.
batch_size
=
batch_size
assert
self
.
batch_size
<
buffer_size
def
reset_state
(
self
):
_MultiProcessZMQDataFlow
.
reset_state
(
self
)
self
.
_guard
=
DataFlowReentrantGuard
()
job_pipe
=
_get_pipe_name
(
"dataflow_MaB_job"
)
result_pipe
=
_get_pipe_name
(
"dataflow_MaB_result"
)
self
.
context
=
zmq
.
Context
()
self
.
socket
=
self
.
context
.
socket
(
zmq
.
PULL
)
self
.
socket
.
set_hwm
(
self
.
buffer_size
*
2
//
self
.
batch_size
)
_bind_guard
(
self
.
socket
,
result_pipe
)
dispatcher
=
MultiProcessMapAndBatchDataZMQ
.
_Dispatcher
(
self
.
ds
,
job_pipe
,
self
.
buffer_size
)
self
.
_proc_ids
=
[
u'{}'
.
format
(
k
)
.
encode
(
'utf-8'
)
for
k
in
range
(
self
.
num_proc
)]
worker_hwm
=
int
(
self
.
buffer_size
*
2
//
self
.
num_proc
)
self
.
_procs
=
[
MultiProcessMapAndBatchDataZMQ
.
_Worker
(
self
.
_proc_ids
[
k
],
self
.
map_func
,
job_pipe
,
result_pipe
,
worker_hwm
,
self
.
batch_size
)
for
k
in
range
(
self
.
num_proc
)]
self
.
_procs
.
append
(
dispatcher
)
self
.
_start_processes
()
def
__iter__
(
self
):
with
self
.
_guard
,
_zmq_catch_error
(
type
(
self
)
.
__name__
):
while
True
:
yield
loads
(
self
.
socket
.
recv
(
copy
=
False
))
def
_pool_map
(
data
):
def
_pool_map
(
data
):
...
@@ -414,6 +525,11 @@ class MultiProcessMapDataComponentSharedArray(DataFlow):
...
@@ -414,6 +525,11 @@ class MultiProcessMapDataComponentSharedArray(DataFlow):
yield
dp
yield
dp
# alias
MultiProcessMapData
=
MultiProcessMapDataZMQ
MultiProcessMapAndBatchData
=
MultiProcessMapAndBatchDataZMQ
if
__name__
==
'__main__'
:
if
__name__
==
'__main__'
:
import
time
import
time
...
...
tensorpack/graph_builder/model_desc.py
View file @
50ff9036
...
@@ -33,12 +33,16 @@ def build_or_reuse_placeholder(tensor_spec):
...
@@ -33,12 +33,16 @@ def build_or_reuse_placeholder(tensor_spec):
assert
"Placeholder"
in
tensor
.
op
.
type
,
"Tensor {} exists but is not a placeholder!"
.
format
(
name
)
assert
"Placeholder"
in
tensor
.
op
.
type
,
"Tensor {} exists but is not a placeholder!"
.
format
(
name
)
assert
tensor_spec
.
is_compatible_with
(
tensor
),
\
assert
tensor_spec
.
is_compatible_with
(
tensor
),
\
"Tensor {} exists but is not compatible with the signature!"
.
format
(
tensor
)
"Tensor {} exists but is not compatible with the signature!"
.
format
(
tensor
)
return
tensor
if
tensor
.
shape
==
tensor_spec
.
shape
:
# It might be desirable to use a placeholder of a different shape in some tower
# (e.g., a less specific shape)
return
tensor
except
KeyError
:
except
KeyError
:
with
tfv1
.
name_scope
(
None
):
# clear any name scope it might get called in
pass
ret
=
tfv1
.
placeholder
(
with
tfv1
.
name_scope
(
None
):
# clear any name scope it might get called in
tensor_spec
.
dtype
,
shape
=
tensor_spec
.
shape
,
name
=
tensor_spec
.
name
)
ret
=
tfv1
.
placeholder
(
return
ret
tensor_spec
.
dtype
,
shape
=
tensor_spec
.
shape
,
name
=
tensor_spec
.
name
)
return
ret
class
InputDesc
(
class
InputDesc
(
...
...
tensorpack/input_source/input_source.py
View file @
50ff9036
...
@@ -454,29 +454,40 @@ class TFDatasetInput(FeedfreeInput):
...
@@ -454,29 +454,40 @@ class TFDatasetInput(FeedfreeInput):
def
__init__
(
self
,
dataset
):
def
__init__
(
self
,
dataset
):
"""
"""
Args:
Args:
dataset (tf.data.Dataset):
dataset (tf.data.Dataset or DataFlow): if a DataFlow, the dataflow
has to be infinite.
"""
"""
if
not
isinstance
(
dataset
,
tf
.
data
.
Dataset
):
if
isinstance
(
dataset
,
tf
.
data
.
Dataset
):
raise
ValueError
(
"TFDatasetInput takes a tf.data.Dataset! Got {}"
.
format
(
dataset
))
self
.
_dataset
=
dataset
self
.
_dataset
=
dataset
self
.
_dataflow
=
None
elif
isinstance
(
dataset
,
DataFlow
):
self
.
_dataset
=
None
self
.
_dataflow
=
dataset
else
:
raise
ValueError
(
"TFDatasetInput takes a tf.data.Dataset or DataFlow! Got {}"
.
format
(
dataset
))
def
_setup
(
self
,
input_signature
):
def
_setup
(
self
,
input_signature
):
self
.
_spec
=
input_signature
self
.
_spec
=
input_signature
types
=
self
.
_dataset
.
output_types
if
self
.
_dataset
is
not
None
:
spec_types
=
tuple
([
k
.
dtype
for
k
in
input_signature
])
types
=
self
.
_dataset
.
output_types
assert
len
(
types
)
==
len
(
spec_types
),
\
spec_types
=
tuple
([
k
.
dtype
for
k
in
input_signature
])
"Dataset and input signature have different length! {} != {}"
.
format
(
assert
len
(
types
)
==
len
(
spec_types
),
\
len
(
types
),
len
(
spec_types
))
"Dataset and input signature have different length! {} != {}"
.
format
(
assert
types
==
spec_types
,
\
len
(
types
),
len
(
spec_types
))
"Data types of dataset and input signature don't match! {} != {}"
.
format
(
assert
types
==
spec_types
,
\
str
(
types
),
str
(
spec_types
))
"Data types of dataset and input signature don't match! {} != {}"
.
format
(
shapes
=
self
.
_dataset
.
output_shapes
str
(
types
),
str
(
spec_types
))
spec_shapes
=
[
k
.
shape
for
k
in
input_signature
]
for
idx
,
(
s1
,
s2
)
in
enumerate
(
zip
(
shapes
,
spec_shapes
)):
shapes
=
self
.
_dataset
.
output_shapes
s2
=
tf
.
TensorShape
(
s2
)
spec_shapes
=
[
k
.
shape
for
k
in
input_signature
]
assert
s2
.
is_compatible_with
(
s1
),
\
for
idx
,
(
s1
,
s2
)
in
enumerate
(
zip
(
shapes
,
spec_shapes
)):
"Input signature '{}' has incompatible shape with dataset! {} vs {}"
.
format
(
s2
=
tf
.
TensorShape
(
s2
)
input_signature
[
idx
]
.
name
,
s2
,
s1
)
assert
s2
.
is_compatible_with
(
s1
),
\
"Input signature '{}' has incompatible shape with dataset! {} vs {}"
.
format
(
input_signature
[
idx
]
.
name
,
s2
,
s1
)
else
:
self
.
_dataset
=
TFDatasetInput
.
dataflow_to_dataset
(
self
.
_dataflow
,
[
x
.
dtype
for
x
in
input_signature
])
self
.
_iterator
=
self
.
_dataset
.
make_initializable_iterator
()
self
.
_iterator
=
self
.
_dataset
.
make_initializable_iterator
()
self
.
_init_op
=
self
.
_iterator
.
initializer
self
.
_init_op
=
self
.
_iterator
.
initializer
...
...
Write
Preview
Markdown
is supported
0%
Try again
or
attach a new file
Attach a file
Cancel
You are about to add
0
people
to the discussion. Proceed with caution.
Finish editing this message first!
Cancel
Please
register
or
sign in
to comment