diff --git a/sdks/python/apache_beam/runners/dask/dask_runner.py b/sdks/python/apache_beam/runners/dask/dask_runner.py index b14449ab2fad..0c800196c2ff 100644 --- a/sdks/python/apache_beam/runners/dask/dask_runner.py +++ b/sdks/python/apache_beam/runners/dask/dask_runner.py @@ -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 @@ -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: @@ -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 @@ -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 diff --git a/sdks/python/apache_beam/runners/dask/dask_runner_test.py b/sdks/python/apache_beam/runners/dask/dask_runner_test.py index e1e5a4403b46..9485068a2ca4 100644 --- a/sdks/python/apache_beam/runners/dask/dask_runner_test.py +++ b/sdks/python/apache_beam/runners/dask/dask_runner_test.py @@ -16,6 +16,8 @@ # import datetime import inspect +import logging +import multiprocessing import typing as t import unittest @@ -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([]) diff --git a/sdks/python/apache_beam/runners/dask/transform_evaluator.py b/sdks/python/apache_beam/runners/dask/transform_evaluator.py index dbf55a3cee0d..23e1bbe2ea44 100644 --- a/sdks/python/apache_beam/runners/dask/transform_evaluator.py +++ b/sdks/python/apache_beam/runners/dask/transform_evaluator.py @@ -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