Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 13 additions & 3 deletions sdks/python/apache_beam/runners/dask/dask_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,13 @@ def _add_argparse_args(cls, parser: argparse.ArgumentParser) -> None:
default=None,
help='The length of each `dask.Bag` partition. When unspecified, '
'an educated guess is made.')
parser.add_argument(
'--dask_lazy_side_inputs',
dest='lazy_side_inputs',
action='store_true',
help='Opt in to partition-at-a-time evaluation of iterable side '
'inputs to reduce peak memory. May increase scheduler overhead and '
'recompute side inputs; AsList still materializes its full view.')


@dataclasses.dataclass
Expand Down Expand Up @@ -166,7 +173,8 @@ def metrics(self):
class DaskRunner(BundleBasedDirectRunner):
"""Executes a pipeline on a Dask distributed client."""
@staticmethod
def to_dask_bag_visitor(bag_kwargs=None) -> PipelineVisitor:
def to_dask_bag_visitor(
bag_kwargs=None, lazy_side_inputs: bool = False) -> PipelineVisitor:
from dask import bag as db

if bag_kwargs is None:
Expand Down Expand Up @@ -210,7 +218,8 @@ def visit_transform(self, transform_node: AppliedPTransform) -> None:
SideInputMap(
type(si),
si._view_options(),
DaskBagWindowedIterator(si_asbag, si._window_mapping_fn)))
DaskBagWindowedIterator(
si_asbag, si._window_mapping_fn, lazy_side_inputs)))

op_kws["side_inputs"] = bag_side_inputs

Expand Down Expand Up @@ -238,11 +247,12 @@ def run_pipeline(self, pipeline, options):
dask_options = options.view_as(DaskOptions).get_all_options(
drop_default=True, current_only=True)
bag_kwargs = DaskOptions._extract_bag_kwargs(dask_options)
lazy_side_inputs = dask_options.pop('lazy_side_inputs', False)
client = ddist.Client(**dask_options)

pipeline.replace_all(dask_overrides())

dask_visitor = self.to_dask_bag_visitor(bag_kwargs)
dask_visitor = self.to_dask_bag_visitor(bag_kwargs, lazy_side_inputs)
pipeline.visit(dask_visitor)
# The dictionary in this visitor keeps a mapping of every Beam
# PTransform to the equivalent Bag operation. This is highly
Expand Down
173 changes: 173 additions & 0 deletions sdks/python/apache_beam/runners/dask/dask_runner_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
#
import datetime
import inspect
import logging
import multiprocessing
import typing as t
import unittest

Expand All @@ -25,17 +27,188 @@
from apache_beam.testing.util import assert_that
from apache_beam.testing.util import equal_to
from apache_beam.transforms import window
from apache_beam.utils.windowed_value import WindowedValue

try:
import dask
import dask.bag as db
import dask.distributed as ddist

from apache_beam.runners.dask.dask_runner import DaskOptions # pylint: disable=ungrouped-imports
from apache_beam.runners.dask.dask_runner import DaskRunner # pylint: disable=ungrouped-imports
from apache_beam.runners.dask.transform_evaluator import DaskBagWindowedIterator # pylint: disable=ungrouped-imports
except (ImportError, ModuleNotFoundError):
raise unittest.SkipTest('Dask must be installed to run tests.')


def _consume_distributed_side_inputs(value, iterable, values, one):
ddist.get_worker() # Assert that side inputs are consumed in a worker task.
return value, list(iterable), values, one


def _run_distributed_side_input_pipeline(lazy, n_workers, threads_per_worker):
with ddist.LocalCluster(n_workers=n_workers,
threads_per_worker=threads_per_worker,
processes=False,
dashboard_address=None,
silence_logs=logging.ERROR) as cluster:
with ddist.Client(cluster):
args = [
'--dask_client_address',
cluster.scheduler_address,
'--dask_partition_size',
'1'
]
if lazy:
args.append('--dask_lazy_side_inputs')
options = PipelineOptions(args)
with test_pipeline.TestPipeline(runner=DaskRunner(),
options=options) as p:
main = p | 'main' >> beam.Create([10])
side = p | 'side' >> beam.Create([2, 3])
singleton = p | 'singleton' >> beam.Create([5])
result = main | beam.Map(
_consume_distributed_side_inputs,
beam.pvalue.AsIter(side),
beam.pvalue.AsList(side),
beam.pvalue.AsSingleton(singleton))
assert_that(result, equal_to([(10, [2, 3], [2, 3], 5)]))
p.result.client.close()


def _run_distributed_partition_error():
@dask.delayed
def fail():
raise ValueError('later partition failed')

def consume():
bag = db.from_delayed([dask.delayed(lambda: [1])(), fail()])
values = iter(
DaskBagWindowedIterator(
bag, window.GlobalWindows(), lazy_side_inputs=True))
assert next(values).value == 1
try:
next(values)
except ValueError as exc:
assert str(exc) == 'later partition failed'
else:
raise AssertionError('later partition error was swallowed')

with ddist.LocalCluster(n_workers=1,
threads_per_worker=1,
processes=False,
dashboard_address=None,
silence_logs=logging.ERROR) as cluster:
with ddist.Client(cluster) as client:
client.submit(consume).result(timeout=20)


class DaskBagWindowedIteratorTest(unittest.TestCase):
def test_default_materializes_bag_before_first_value(self):
computed = []

@dask.delayed
def partition(number):
computed.append(number)
return [number]

bag = db.from_delayed([partition(1), partition(2)])
with dask.config.set(scheduler='synchronous'):
values = iter(DaskBagWindowedIterator(bag, window.GlobalWindows()))
self.assertEqual(next(values).value, 1)
self.assertCountEqual(computed, [1, 2])

def test_computes_partitions_as_values_are_consumed(self):
computed = []

@dask.delayed
def partition(number):
computed.append(number)
return [number]

bag = db.from_delayed([partition(1), partition(2)])
with dask.config.set(scheduler='synchronous'):
values = iter(
DaskBagWindowedIterator(
bag, window.GlobalWindows(), lazy_side_inputs=True))
self.assertEqual(next(values).value, 1)
self.assertEqual(computed, [1])
self.assertEqual(next(values).value, 2)
self.assertEqual(computed, [1, 2])
self.assertRaises(StopIteration, next, values)

def test_empty_partitions_and_repeated_consumption(self):
bag = db.from_delayed([
dask.delayed(lambda: [])(),
dask.delayed(lambda: [1, 2])(),
])
side_input = DaskBagWindowedIterator(
bag, window.GlobalWindows(), lazy_side_inputs=True)
with dask.config.set(scheduler='synchronous'):
self.assertEqual([value.value for value in side_input], [1, 2])
self.assertEqual([value.value for value in side_input], [1, 2])

def test_partition_error_is_raised_when_reached(self):
@dask.delayed
def failing_partition():
raise ValueError('partition failed')

bag = db.from_delayed([
dask.delayed(lambda: [1])(),
failing_partition(),
])
with dask.config.set(scheduler='synchronous'):
values = iter(
DaskBagWindowedIterator(
bag, window.GlobalWindows(), lazy_side_inputs=True))
self.assertEqual(next(values).value, 1)
with self.assertRaisesRegex(ValueError, 'partition failed'):
next(values)

def test_preserves_order_and_window_conversion(self):
existing = WindowedValue('existing', 1, (window.IntervalWindow(0, 5), ))
bag = db.from_sequence(
[window.TimestampedValue('timestamped', 7), existing, 'plain'],
partition_size=1)
with dask.config.set(scheduler='synchronous'):
values = list(
DaskBagWindowedIterator(
bag, window.FixedWindows(5), lazy_side_inputs=True))
self.assertEqual([value.value for value in values],
['timestamped', 'existing', 'plain'])
self.assertEqual(values[0].windows, (window.IntervalWindow(5, 10), ))
self.assertIs(values[1], existing)
self.assertEqual(values[2].windows, (window.GlobalWindow(), ))


class DaskDistributedSideInputTest(unittest.TestCase):
def _run_with_watchdog(self, target, *args):
process = multiprocessing.get_context('spawn').Process(
target=target, args=args)
process.start()
process.join(timeout=45)
if process.is_alive():
process.terminate()
process.join(timeout=5)
if process.is_alive():
process.kill()
process.join()
self.fail('distributed side-input computation timed out')
self.assertEqual(process.exitcode, 0)

def test_lazy_pipeline_with_one_worker_one_thread(self):
self._run_with_watchdog(_run_distributed_side_input_pipeline, True, 1, 1)

def test_lazy_pipeline_with_two_workers(self):
self._run_with_watchdog(_run_distributed_side_input_pipeline, True, 2, 2)

def test_default_pipeline_with_one_worker_one_thread(self):
self._run_with_watchdog(_run_distributed_side_input_pipeline, False, 1, 1)

def test_later_partition_error_in_worker(self):
self._run_with_watchdog(_run_distributed_partition_error)


class DaskOptionsTest(unittest.TestCase):
def test_parses_connection_timeout__defaults_to_none(self):
default_options = PipelineOptions([])
Expand Down
42 changes: 36 additions & 6 deletions sdks/python/apache_beam/runners/dask/transform_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,17 +83,47 @@ def defenestrate(x):

@dataclasses.dataclass
class DaskBagWindowedIterator:
"""Iterator for `apache_beam.transforms.sideinputs.SideInputMap`"""
"""Iterator for `apache_beam.transforms.sideinputs.SideInputMap`.

The default computes the whole bag for compatibility. With lazy side inputs,
each partition is computed when reached by iteration. Computations launched
inside a Dask worker use worker_client to release the worker's task slot.
"""

bag: db.Bag
window_fn: WindowFn
lazy_side_inputs: bool = False

@staticmethod
def _compute(collection):
from dask.distributed import get_worker
from dask.distributed import worker_client

try:
get_worker()
except ValueError:
# Outside a worker, use the configured Dask scheduler.
return collection.compute()

# A blocking nested computation must release the worker's task slot.
# Close the context before yielding so partial iteration cannot leave a
# worker thread seceded from its pool.
with worker_client() as client:
return client.compute(collection).result()

def __iter__(self):
# FIXME(cisaacstern): list() is likely inefficient, since it presumably
# materializes the full result before iterating over it. doing this for
# now as a proof-of-concept. can we can generate results incrementally?
for result in list(self.bag):
yield get_windowed_value(result, self.window_fn)
if self.lazy_side_inputs:
# AsIter can consume one partition at a time. AsList still materializes
# the complete view, as required by its side-input semantics.
for partition in self.bag.to_delayed():
for result in self._compute(partition):
yield get_windowed_value(result, self.window_fn)
else:
# FIXME(cisaacstern): The original list(self.bag) materializes the full
# side input before iteration. Keep that behavior by default until a
# shared or storage-backed side-input design can avoid its memory cost.
for result in self._compute(self.bag):
yield get_windowed_value(result, self.window_fn)


@dataclasses.dataclass
Expand Down
Loading