Skip to content
Closed
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
111 changes: 111 additions & 0 deletions sdks/python/apache_beam/transforms/partitioners.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
import uuid
import apache_beam as beam
from apache_beam import pvalue
from typing import Optional, TypeVar
from typing import Tuple
from typing import Any
from typing import Callable

T = TypeVar('T')


class Top(beam.PTransform):
"""
A PTransform that takes a PCollection and partitions it into two
PCollections. The first PCollection contains the largest n elements of the
input PCollection, and the second PCollection contains the remaining
elements of the input PCollection.

Parameters:
n: The number of elements to take from the input PCollection.
key: A function that takes an element of the input PCollection and
returns a value to compare for the purpose of determining the top n
elements, similar to Python's built-in sorted function.
reverse: If True, the top n elements will be the n smallest elements of
the input PCollection.

Example usage:

>>> with beam.Pipeline() as p:
... top, remaining = (p
... | beam.Create(list(range(10)))
... | partitioners.Top(3))
... # top will contain [7, 8, 9]
... # remaining will contain [0, 1, 2, 3, 4, 5, 6]

.. note::

This transform requires that the top PCollection fit into memory.

"""
def __init__(
self, n: int, key: Optional[Callable[[Any], Any]] = None, reverse=False):
_validate_nonzero_positive_int(n)
self.n = n
self.key = key
self.reverse = reverse

def expand(self,
pcoll) -> Tuple[pvalue.PCollection[T], pvalue.PCollection[T]]:
# **Illustrative Example:**
# Our goal is to return two pcollections, top and
# remaining.

# Suppose you want to take the top element from `[1, 2, 2]`. Since we have
# identical elements, we need to be able to uniquely identify each one,
# so we assign a unique ID to each:
# `inputs_with_ids: [(1, "A"), (2, "B"), (2, "C")]`

# Then we sample, e.g.
# ``` sample: [(2, "B")] ```
# To get our goal `top` pcollection, we just strip the uuids from
# that sample.

# Now to get the `top` pcollection, we need to return essentially
# `inputs_with_ids` but without any of the elements fom the sample. To
# do this, we create a set from `sample`, getting `sample_ids:
# [set("B")]`. Now that we have this set, we can create our
# `remaining_with_ids` pcollection by just filtering out
# `inputs_with_ids` and checking for each element "Does this element's
# corresponding ID exist in `sample_ids`?"

# Finally, we just return `top` and strip the IDs as we no longer
# need them and the user doesn't care about them.
wrapped_key = lambda elem: self.key(elem[0]) if self.key else elem[0]
inputs_with_ids = (pcoll | beam.Map(_add_uuid))
sample = (
inputs_with_ids
| beam.combiners.Top.Of(self.n, key=wrapped_key, reverse=self.reverse))
sample_ids = (
sample
| beam.Map(lambda sample_list: set(ele[1] for ele in sample_list)))

def elem_is_not_sampled(elem, sampled_set):
return elem[1] not in sampled_set

remaining = (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we can both simplify this implementation and significantly improve performance using user state. Basically the idea would be to use bag state and:

  1. Add a dummy key if the input is unkeyed
  2. Read state to get our current running set of N elements.
  3. Scan that state to get the smallest element
  4. Check if the current element is larger than the smallest element grabbed from state. If yes, then add it to state. If no, emit it into the non-top PCollection
  5. In StartBundle, set an event time timer that expires at the end of the window. In that timer, if you have any buffered items in state emit them and clear your state.
  6. Remove the dummy key

This would have a few advantages:

  1. It wouldn't block execution of later steps for the non-Top PCollection. So, you could partition your data and wait until the end of the window to process your Top PCollection, but you could start processing elements in your non-top PCollection as soon as you've seen enough data to guarantee that they can't be a part of your Top PCollection.
  2. It wouldn't require doing a join which could be at least a bit expensive.
  3. You get per-key Top for free since state is per-key - you just would skip adding/removing the dummy key.

There's even a few optimizations we could potentially add:

  1. Store the current minimum value in valueState (duplicated from the bag state) and only read that initially (since it is a potentially much cheaper read if N gets large)
  2. Cache the current minimum value in the DoFn and route any future entries that are smaller than that to the non-Top PCollection. This would get out of date, so we'd still need to do the comparison to state if our value was larger than the min_val.

Thoughts?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the review!

That does seem simple. I don't have a strong understanding of how user state works. Does this imply that a single worker will need to go through the entire set of inputs? Won't that become a bottleneck or am I misunderstanding how state works?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It does, though in exchange you're getting the ability to start processing your data in the rest of your pipeline before you reach the end of the window (which probably saves time/resources in most use cases, especially in streaming). I think this is a fair concern though, especially in batch; one optimization we could make to greatly reduce our bottleneck in cases where Top is small would be to add a pre-step where we filter out the non-Top elements and send them downstream to be joined with the remaining non-Top elements using a Flatten (something like _TopPerBundle). Cases where Top is large will still end up running into a similar bottleneck even if we use a combiner since we won't be able to do much combiner lifting (local reduction before sending it over the wire).

I'll also note that state is per key/window, so we're actually just talking the set of inputs for a single window and you will get parallelization with multiple windows.

@hjtran hjtran Oct 26, 2023

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is batch processing rare enough that it's not considered in "most use cases"? (Coming from a batch-only user).

I don't quite follow your suggestion.

one optimization we could make to greatly reduce our bottleneck in cases where Top is small would be to add a pre-step where we filter out the non-Top elements

Isn't this still doing the join that's being proposed in the current changeset?

send them downstream to be joined with the remaining non-Top elements

I don't follow the distinctino between these "remaining non-Top elements" and the non-Top elements we're filtering for in the pre-step

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is batch processing rare enough that it's not considered in "most use cases"? (Coming from a batch-only user).

No, in general I think we have a pretty even spread of batch/streaming usage (though its hard to say for sure). My point was more that I think this is beneficial for most use cases and this is especially true in streaming mode (but also many and possibly most batch cases). Even in batch mode, everything I said is true, you're just also probably more likely to have scenarios where you're particularly concerned about quick fan out. I do think my suggested optimization (explained in more detail below) resolves most of these concerns though.

Basically, we're making a tradeoff, but I think the large majority of the time we can come out ahead with a stateful approach.

Isn't this still doing the join that's being proposed in the current changeset?

No, there shouldn't be a join (Flatten is the closest thing, but that's not actually a join, its just conceptually treating 2 PCollections as one). Let me try explaining again in more depth; basically your flow would be:

top_candidates, non_top1 = (pcoll
| _TopPerBundle(...)) # modified to emit 2 pcollections
top, non_top2 = (top_candidates
| StatefulTopPartitioner(...)) # likely a composite where you add a key as discussed above
non_top = Flatten(non_top1, non_top2)

return top, non_top

In this way, you don't need to do any expensive joins, you get the benefit of parallelism for _TopPerBundle (which should reduce your cardinality by O(1000s)/N where N is your Top size in batch mode (it won't help much in streaming which usually has smaller bundles depending on the runner), and your bottleneck is on a small subset of the data.

_TopPerBundle would emit 2 pcollections, the first containing the top N elements from each bundle, and the 2nd containing the remaining elements.

StatefulTopPartitioner would do the same, but over the entire input dataset.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ah that makes sense! I'll take a shot at implementing that, thanks for the explanation!

inputs_with_ids
| beam.Filter(
elem_is_not_sampled,
sampled_set=beam.pvalue.AsSingleton(sample_ids))
| beam.Map(_strip_uuid))
sample_as_pcoll = (
sample
| beam.FlatMap(lambda x: x)
| "StripSampleIDs" >> beam.Map(_strip_uuid))
return sample_as_pcoll, remaining


def _validate_nonzero_positive_int(n: Optional[Any]) -> None:
if not isinstance(n, int):
raise ValueError("n must be an int")
if n <= 0:
raise ValueError("n must be a positive int")


def _add_uuid(element: T) -> Tuple[T, str]:
return element, uuid.uuid4().hex


def _strip_uuid(element: Tuple[T, str]) -> T:
return element[0]
95 changes: 95 additions & 0 deletions sdks/python/apache_beam/transforms/partitioners_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
import doctest
import pytest
import unittest

import apache_beam as beam
from apache_beam.testing.util import assert_that
from apache_beam.testing.util import equal_to
from apache_beam.transforms import partitioners
from apache_beam.testing.test_pipeline import TestPipeline


class TopTest(unittest.TestCase):
def test_bad_n(self):
with pytest.raises(ValueError):
partitioners.Top(0)
with pytest.raises(ValueError):
partitioners.Top(-1)
with pytest.raises(ValueError):
partitioners.Top(1.5)

def test_empty(self):
with TestPipeline() as p:
sample, remaining = (p
| beam.Create([], reshuffle=False)
| partitioners.Top(3))
assert_that(
sample | "CountSample" >> beam.combiners.Count.Globally(),
equal_to([0]),
label="assert0")
assert_that(
remaining | "CountRemaining" >> beam.combiners.Count.Globally(),
equal_to([0]),
label="assert1")

def test_Top(self):

with TestPipeline() as p:
sample, remaining = (p
| beam.Create([1, 1, 2, 2, 3, 3], reshuffle=False)
| partitioners.Top(3))
assert_that(
sample | "SampleAsList" >> beam.combiners.ToList()
| "SortSample" >> beam.Map(sorted),
equal_to([[
2,
3,
3,
]]),
label="assert0")
assert_that(
remaining | "RemainingAsList" >> beam.combiners.ToList()
| "SortRemaining" >> beam.Map(sorted),
equal_to([[1, 1, 2]]),
label="assert1")

def test_Top_key(self):

with TestPipeline() as p:
sample, remaining = (p
| beam.Create([1, 1, 2, 2, 3, 3],
reshuffle=False)
| partitioners.Top(3, key=lambda x: -x))
assert_that(
sample | "SampleAsList" >> beam.combiners.ToList()
| "SortSample" >> beam.Map(sorted),
equal_to([[1, 1, 2]]),
label="assert0")
assert_that(
remaining | "RemainingAsList" >> beam.combiners.ToList()
| "SortRemaining" >> beam.Map(sorted),
equal_to([[2, 3, 3]]),
label="assert1")

def test_Top_reverse(self):

with TestPipeline() as p:
sample, remaining = (p
| beam.Create([1, 1, 2, 2, 3, 3],
reshuffle=False)
| partitioners.Top(3, reverse=True))
assert_that(
sample | "SampleAsList" >> beam.combiners.ToList()
| "SortSample" >> beam.Map(sorted),
equal_to([[1, 1, 2]]),
label="assert0")
assert_that(
remaining | "RemainingAsList" >> beam.combiners.ToList()
| "SortRemaining" >> beam.Map(sorted),
equal_to([[2, 3, 3]]),
label="assert1")


class DocTest(unittest.TestCase):
def test_docs(self):
doctest.testmod(partitioners)