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
152 changes: 108 additions & 44 deletions sorts/odd_even_transposition_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,14 @@
This implementation represents each variable in the list with a process and
each process communicates with its neighboring processes in the list to perform
comparisons.
They are synchronized with locks and message passing but other forms of
They are synchronized with message passing but other forms of
synchronization could be used.
"""

import multiprocessing as mp

# lock used to ensure that two processes do not access a pipe at the same time
# NOTE This breaks testing on build runner. May work better locally
# process_lock = mp.Lock()
from multiprocessing.connection import Connection, wait
from multiprocessing.reduction import ForkingPickler
from typing import Any, Protocol

"""
The function run by the processes that sorts the list
Expand All @@ -29,46 +28,56 @@
"""


def oe_process(
position,
value,
l_send,
r_send,
lr_cv,
rr_cv,
result_pipe,
multiprocessing_context,
) -> None:
process_lock = multiprocessing_context.Lock()
class Comparable(Protocol):
def __lt__(self, other: Any, /) -> bool: ...


def oe_process[T: Comparable](
position: int,
value: T,
l_send: tuple[Connection, Connection] | None,
r_send: tuple[Connection, Connection] | None,
lr_cv: tuple[Connection, Connection] | None,
rr_cv: tuple[Connection, Connection] | None,
result_pipe: tuple[Connection, Connection],
) -> None:
# we perform n swaps since after n swaps we know we are sorted
# we *could* stop early if we are sorted already, but it takes as long to
# find out we are sorted as it does to sort the list with this algorithm
for i in range(10):
if (i + position) % 2 == 0 and r_send is not None:
# send your value to your right neighbor
with process_lock:
try:
for i in range(10):
if (i + position) % 2 == 0 and r_send is not None and rr_cv is not None:
# send your value to your right neighbor
r_send[1].send(value)

# receive your right neighbor's value
with process_lock:
# receive your right neighbor's value
temp = rr_cv[0].recv()

# take the lower value since you are on the left
value = min(value, temp)
elif (i + position) % 2 != 0 and l_send is not None:
# send your value to your left neighbor
with process_lock:
# take the lower value since you are on the left
value = temp if temp < value else value
elif (i + position) % 2 != 0 and l_send is not None and lr_cv is not None:
# send your value to your left neighbor
l_send[1].send(value)

# receive your left neighbor's value
with process_lock:
# receive your left neighbor's value
temp = lr_cv[0].recv()

# take the higher value since you are on the right
value = max(value, temp)
# after all swaps are performed, send the values back to main
result_pipe[1].send(value)
# take the higher value since you are on the right
value = temp if value < temp else value
# after all swaps are performed, send the values back to main
result_pipe[1].send((value, None))
except Exception as error: # noqa: BLE001 -- propagate worker errors to the caller
try:
payload = ForkingPickler.dumps((None, error))
except Exception: # noqa: BLE001 -- user exceptions can fail to pickle
fallback = RuntimeError(f"{type(error).__name__}: {error}")
payload = ForkingPickler.dumps((None, fallback))
result_pipe[1].send_bytes(payload)
finally:
for pipe in (l_send, r_send, lr_cv, rr_cv, result_pipe):
if pipe is not None:
for connection in pipe:
connection.close()


"""
Expand All @@ -78,8 +87,16 @@ def oe_process(
"""


def odd_even_transposition(arr):
def odd_even_transposition[T: Comparable](arr: list[T]) -> list[T]:
"""
Sort in place, propagating worker errors to the caller. Unpickleable
exceptions become RuntimeError with the original type name and message.

>>> odd_even_transposition([])
[]
>>> values = [42]
>>> odd_even_transposition(values) is values
True
>>> odd_even_transposition(list(range(10)[::-1])) == sorted(list(range(10)[::-1]))
True
>>> odd_even_transposition(["a", "x", "c"]) == sorted(["x", "a", "c"])
Expand All @@ -98,7 +115,21 @@ def odd_even_transposition(arr):
>>> unsorted_list = [-442, -98, -554, 266, -491, 985, -53, -529, 82, -429]
>>> odd_even_transposition(unsorted_list) == sorted(unsorted_list + [1])
False
>>> values = ["c", "a", "b"]
>>> odd_even_transposition(values) is values
True
>>> values
['a', 'b', 'c']
>>> odd_even_transposition([2.5, -1, 0.0])
[-1, 0.0, 2.5]
>>> odd_even_transposition([1, "a"]) # doctest: +IGNORE_EXCEPTION_DETAIL
Traceback (most recent call last):
...
TypeError: '<' not supported between instances of 'str' and 'int'
"""
if len(arr) < 2:
return arr

# spawn method is considered safer than fork
multiprocessing_context = mp.get_context("spawn")

Expand All @@ -112,6 +143,7 @@ def odd_even_transposition(arr):
# of the loop
temp_rs = multiprocessing_context.Pipe()
temp_rr = multiprocessing_context.Pipe()
neighbor_pipes = [temp_rs, temp_rr]
process_array_.append(
multiprocessing_context.Process(
target=oe_process,
Expand All @@ -123,7 +155,6 @@ def odd_even_transposition(arr):
None,
temp_rr,
result_pipe[0],
multiprocessing_context,
),
)
)
Expand All @@ -133,6 +164,7 @@ def odd_even_transposition(arr):
for i in range(1, len(arr) - 1):
temp_rs = multiprocessing_context.Pipe()
temp_rr = multiprocessing_context.Pipe()
neighbor_pipes.extend((temp_rs, temp_rr))
process_array_.append(
multiprocessing_context.Process(
target=oe_process,
Expand All @@ -144,7 +176,6 @@ def odd_even_transposition(arr):
temp_lr,
temp_rr,
result_pipe[i],
multiprocessing_context,
),
)
)
Expand All @@ -162,19 +193,52 @@ def odd_even_transposition(arr):
temp_lr,
None,
result_pipe[len(arr) - 1],
multiprocessing_context,
),
)
)

# start the processes
for p in process_array_:
p.start()
started_processes = []
try:
for process in process_array_:
process.start()
started_processes.append(process)

pending = {pipe[0]: position for position, pipe in enumerate(result_pipe)}
sentinels = {
process.sentinel: position
for position, process in enumerate(process_array_)
}
Comment thread
Ethereal49 marked this conversation as resolved.
values = list(arr)
while pending:
ready = set(wait([*pending, *sentinels]))
for connection in pending.keys() & ready:
position = pending.pop(connection)
value, error = connection.recv()
if error is not None:
raise error
values[position] = value
for process_sentinel in sentinels.keys() & ready:
position = sentinels.pop(process_sentinel)
connection = result_pipe[position][0]
if connection in pending and not connection.poll():
raise RuntimeError(
"Sorting worker exited without returning a result"
)

# wait for the processes to end and write their values to the list
for p in range(len(result_pipe)):
arr[p] = result_pipe[p][0].recv()
process_array_[p].join()
# Do not partially overwrite the input if another worker fails.
arr[:] = values
except BaseException:
for process in started_processes:
if process.is_alive():
process.terminate()
raise
finally:
for process in started_processes:
process.join()
process.close()
for pipe in result_pipe + neighbor_pipes:
for connection in pipe:
connection.close()
return arr


Expand Down
84 changes: 84 additions & 0 deletions tests/test_sorts.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,12 @@
separately below.
"""

import multiprocessing as mp
import os
import signal
from dataclasses import dataclass
from typing import NamedTuple
from unittest.mock import patch

import pytest

Expand All @@ -41,6 +45,9 @@
from sorts.merge_insertion_sort import merge_insertion_sort
from sorts.merge_sort import merge_sort
from sorts.odd_even_sort import odd_even_sort
from sorts.odd_even_transposition_parallel import (
odd_even_transposition as parallel_odd_even_transposition,
)
from sorts.odd_even_transposition_single_threaded import odd_even_transposition
from sorts.pancake_sort import pancake_sort
from sorts.patience_sort import patience_sort
Expand Down Expand Up @@ -244,3 +251,80 @@ def test_bitonic_sort_comparable_items() -> None:

with pytest.raises(TypeError):
bitonic_sort([1, "two", 3, "four"], 0, 4, 1)


@dataclass
class UnpicklableComparison:
value: int

def __lt__(self, _other: object, /) -> bool:
class LocalComparisonError(TypeError):
pass

raise LocalComparisonError("comparison failed")


def _check_parallel_odd_even_transposition(
case: list[object], error: type[Exception] | None
) -> None:
# Give this probe and its workers a process group that the test alone owns.
os.setsid()
collection = list(case)
if error is not None:
with pytest.raises(error) as caught:
parallel_odd_even_transposition(collection)
if isinstance(case[0], UnpicklableComparison):
assert "LocalComparisonError: comparison failed" in str(caught.value)
assert collection == case
elif len(collection) < 2:
with patch(
"multiprocessing.process.BaseProcess.start",
side_effect=AssertionError("trivial input must not start a worker"),
):
assert parallel_odd_even_transposition(collection) is collection
assert all(item is original for item, original in zip(collection, case))
else:
assert parallel_odd_even_transposition(collection) is collection
assert collection == sorted(case)
assert not mp.active_children()


@pytest.mark.skipif(
os.name != "posix", reason="timeout cleanup requires process groups"
)
@pytest.mark.parametrize(
("case", "error"),
[
([], None),
([1], None),
([Person()], None),
(["c", "a", "b"], None),
([2.5, -1, 0.0], None),
([Person(cost=100.0), Person(cost=-100.0), Person(name="Al")], None),
([Dog(weight=15.5), Dog(weight=15.1), Dog(name="Buddy")], None),
([1, "a"], TypeError),
([3, 2, "a", 1], TypeError),
([UnpicklableComparison(3), UnpicklableComparison(2)], RuntimeError),
],
)
def test_parallel_odd_even_transposition(
case: list[object], error: type[Exception] | None
) -> None:
process = mp.get_context("spawn").Process(
target=_check_parallel_odd_even_transposition, args=(case, error)
)
process.start()
try:
process.join(timeout=10)
assert not process.is_alive(), (
"parallel sorting did not finish within 10 seconds"
)
assert process.exitcode == 0
finally:
if process.is_alive():
try:
os.killpg(process.pid, signal.SIGKILL)
except ProcessLookupError:
process.kill()
process.join()
process.close()
Loading