diff --git a/cpp/tests/CMakeLists.txt b/cpp/tests/CMakeLists.txt index 12421123a..4e7278a59 100644 --- a/cpp/tests/CMakeLists.txt +++ b/cpp/tests/CMakeLists.txt @@ -103,6 +103,7 @@ if(RAPIDSMPF_HAVE_STREAMING) streaming/test_channel.cpp streaming/test_error_handling.cpp streaming/test_fanout.cpp + streaming/test_leaf_actor.cpp streaming/test_lineariser.cpp streaming/test_memory_reserve_or_wait.cpp streaming/test_message.cpp @@ -123,8 +124,7 @@ if(BUILD_CUDF_TESTS) if(RAPIDSMPF_HAVE_STREAMING) target_sources( - test_sources PRIVATE streaming/test_allgather.cpp streaming/test_leaf_actor.cpp - streaming/test_shuffler.cpp + test_sources PRIVATE streaming/test_allgather.cpp streaming/test_shuffler.cpp ) endif() endif() diff --git a/cpp/tests/streaming/test_leaf_actor.cpp b/cpp/tests/streaming/test_leaf_actor.cpp index cd5e90290..481644b68 100644 --- a/cpp/tests/streaming/test_leaf_actor.cpp +++ b/cpp/tests/streaming/test_leaf_actor.cpp @@ -4,17 +4,15 @@ */ #include +#include #include +#include +#include #include #include #include -#include -#include - -#include -#include #include #include #include @@ -23,7 +21,6 @@ #include #include -#include "../utils.hpp" #include "base_streaming_fixture.hpp" using namespace rapidsmpf; @@ -33,12 +30,11 @@ namespace actor = rapidsmpf::streaming::actor; using StreamingLeafTasks = BaseStreamingFixture; TEST_F(StreamingLeafTasks, PushAndPullChunks) { - constexpr int num_rows = 100; constexpr int num_chunks = 10; - std::vector expects; + std::vector expects; for (int i = 0; i < num_chunks; ++i) { - expects.emplace_back(random_table_with_index(i, num_rows, 0, 10)); + expects.push_back(i * 10); } std::vector actors; @@ -49,15 +45,7 @@ TEST_F(StreamingLeafTasks, PushAndPullChunks) { std::vector inputs; for (int i = 0; i < num_chunks; ++i) { inputs.emplace_back( - cudf_streaming::streaming::to_message( - i, - std::make_unique( - std::make_unique( - expects[i], stream, ctx->br()->device_mr() - ), - stream - ) - ) + i, std::make_unique(expects[i]), ContentDescription{} ); } @@ -72,10 +60,7 @@ TEST_F(StreamingLeafTasks, PushAndPullChunks) { EXPECT_EQ(expects.size(), outputs.size()); for (std::size_t i = 0; i < expects.size(); ++i) { EXPECT_EQ(outputs[i].sequence_number(), i); - CUDF_TEST_EXPECT_TABLES_EQUIVALENT( - outputs[i].get().table_view(), - expects[i].view() - ); + EXPECT_EQ(outputs[i].release(), expects[i]); } } diff --git a/python/rapidsmpf/rapidsmpf/tests/streaming/test_allgather.py b/python/rapidsmpf/rapidsmpf/tests/streaming/test_allgather.py index d8d6ed019..447e5dcc7 100644 --- a/python/rapidsmpf/rapidsmpf/tests/streaming/test_allgather.py +++ b/python/rapidsmpf/rapidsmpf/tests/streaming/test_allgather.py @@ -6,23 +6,14 @@ from contextlib import nullcontext from typing import TYPE_CHECKING -import numpy as np -import pylibcudf as plc import pytest -pytest.importorskip("cudf_streaming") -from cudf_streaming.integrations.partition import unpack_and_concat -from cudf_streaming.streaming.table_chunk import TableChunk - -from rapidsmpf.memory.packed_data import PackedData from rapidsmpf.streaming.chunks.packed_data import PackedDataChunk from rapidsmpf.streaming.coll.allgather import AllGather, allgather from rapidsmpf.streaming.core.actor import define_actor, run_actor_network from rapidsmpf.streaming.core.leaf_actor import pull_from_channel, push_to_channel from rapidsmpf.streaming.core.message import Message -from rapidsmpf.testing import assert_eq - -cudf = pytest.importorskip("cudf") +from rapidsmpf.testing import generate_packed_data, validate_packed_data if TYPE_CHECKING: from collections.abc import Awaitable @@ -33,34 +24,28 @@ from rapidsmpf.streaming.core.context import Context +def _make_chunk(context: Context, num_rows: int, offset: int) -> PackedDataChunk: + stream = context.get_stream_from_pool() + return PackedDataChunk.from_packed_data( + generate_packed_data(num_rows, offset, stream, context.br()), + br=context.br(), + ) + + +def _validate_message( + context: Context, msg: Message[PackedDataChunk], num_rows: int, offset: int +) -> None: + packed = PackedDataChunk.from_message(msg, br=context.br()).to_packed_data() + validate_packed_data(packed, num_rows, offset) + + def test_allgather_actor(context: Context, comm: Communicator) -> None: if comm.nranks != 1: pytest.skip("Only support single-rank runs") num_rows = 1000 op_id = 0 - stream = context.get_stream_from_pool() - input_tables = [ - plc.Table( - [ - plc.Column.from_array( - np.arange(num_rows, dtype=np.int32) + i * num_rows, stream=stream - ) - ] - ) - for i in range(3) - ] - inputs = [ - PackedDataChunk.from_packed_data( - PackedData.from_cudf_packed_columns( # type: ignore[attr-defined] - plc.contiguous_split.pack(table, stream=stream), - stream, - context.br(), - ), - br=context.br(), - ) - for table in input_tables - ] + inputs = [_make_chunk(context, num_rows, i * num_rows) for i in range(3)] actors = [] ch1: Channel[PackedDataChunk] = context.create_channel() @@ -77,18 +62,11 @@ def test_allgather_actor(context: Context, comm: Communicator) -> None: actors.append(actor) run_actor_network(context, actors=actors) - result = unpack_and_concat( - ( - PackedDataChunk.from_message(msg, br=context.br()).to_packed_data() - for msg in deferred.release() - ), - stream, - context.br(), - ) - - expect = plc.concatenate.concatenate(input_tables, stream=stream) - stream.synchronize() - assert_eq(result, expect) + results = deferred.release() + assert len(results) == len(inputs) + for i, msg in enumerate(results): + assert msg.sequence_number == i + _validate_message(context, msg, num_rows, i * num_rows) @define_actor() @@ -96,35 +74,16 @@ async def generate_inputs( context: Context, ch: Channel[PackedDataChunk], num_rows: int, num_chunks: int ) -> None: for i in range(num_chunks): - stream = context.get_stream_from_pool() - table = plc.Table( - [ - plc.Column.from_array( - np.arange(num_rows, dtype=np.int32) + i * num_rows, stream=stream - ) - ] - ) - msg = Message( - i, - PackedDataChunk.from_packed_data( - PackedData.from_cudf_packed_columns( # type: ignore[attr-defined] - plc.contiguous_split.pack(table, stream=stream), - stream, - context.br(), - ), - br=context.br(), - ), - ) - await ch.send(context, msg) + await ch.send(context, Message(i, _make_chunk(context, num_rows, i * num_rows))) await ch.drain(context) @define_actor() -async def allgather_and_concat( +async def allgather_and_forward( context: Context, comm: Communicator, ch_in: Channel[PackedDataChunk], - ch_out: Channel[TableChunk], + ch_out: Channel[PackedDataChunk], op_id: int, use_context_manager: bool, # noqa: FBT001 ) -> None: @@ -137,12 +96,14 @@ async def allgather_and_concat( if not use_context_manager: gather.insert_finished() gathered = await gather.extract_all(context, ordered=True) - stream = context.get_stream_from_pool() - table = unpack_and_concat(gathered, stream, context.br()) - to_send = TableChunk.from_pylibcudf_table( - table, stream, exclusive_view=True, br=context.br() - ) - await ch_out.send(context, Message(0, to_send)) + for sequence, packed in enumerate(gathered): + await ch_out.send( + context, + Message( + sequence, + PackedDataChunk.from_packed_data(packed, br=context.br()), + ), + ) await ch_out.drain(context) @@ -158,29 +119,22 @@ def test_allgather_object_interface( pytest.skip("Only support single-rank runs") ch_in: Channel[PackedDataChunk] = context.create_channel() - ch_out: Channel[TableChunk] = context.create_channel() + ch_out: Channel[PackedDataChunk] = context.create_channel() actors: list[CppActor | Awaitable[None]] = [] num_rows = 100 num_chunks = 10 op_id = 0 actors.append(generate_inputs(context, ch_in, num_rows, num_chunks)) actors.append( - allgather_and_concat(context, comm, ch_in, ch_out, op_id, use_context_manager) + allgather_and_forward(context, comm, ch_in, ch_out, op_id, use_context_manager) ) actor, deferred = pull_from_channel(context, ch_out) actors.append(actor) run_actor_network(context, actors=actors) - (result_msg,) = deferred.release() - result = TableChunk.from_message(result_msg, br=context.br()) - expect = plc.Table( - [ - plc.Column.from_array( - np.arange(num_rows * num_chunks, dtype=np.int32), stream=result.stream - ) - ] - ) - got = result.table_view() - result.stream.synchronize() - assert_eq(expect, got) + results = deferred.release() + assert len(results) == num_chunks + for i, msg in enumerate(results): + assert msg.sequence_number == i + _validate_message(context, msg, num_rows, i * num_rows) diff --git a/python/rapidsmpf/rapidsmpf/tests/streaming/test_define_actor.py b/python/rapidsmpf/rapidsmpf/tests/streaming/test_define_actor.py index 955fc0fb8..55e67a7d7 100644 --- a/python/rapidsmpf/rapidsmpf/tests/streaming/test_define_actor.py +++ b/python/rapidsmpf/rapidsmpf/tests/streaming/test_define_actor.py @@ -5,64 +5,35 @@ from typing import TYPE_CHECKING -import pylibcudf as plc import pytest -pytest.importorskip("cudf_streaming") -from cudf_streaming.streaming.table_chunk import TableChunk - from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk from rapidsmpf.streaming.core.actor import define_actor, run_actor_network from rapidsmpf.streaming.core.leaf_actor import pull_from_channel, push_to_channel from rapidsmpf.streaming.core.message import Message -from rapidsmpf.testing import assert_eq - -cudf = pytest.importorskip("cudf") @pytest.fixture -def expects() -> list[plc.Table]: - return [ - plc.Table( - [ - plc.Column.from_iterable_of_py( - [1 * seq, 2 * seq, 3 * seq], plc.DataType(plc.TypeId.INT64) - ) - ] - ) - for seq in range(10) - ] +def expects() -> list[tuple[int, int, int]]: + return [(seq, seq * 2, seq * 3) for seq in range(10)] if TYPE_CHECKING: - from rmm.pylibrmm.stream import Stream - from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context -def test_send_table_chunks( - context: Context, stream: Stream, expects: list[plc.Table] +def test_send_arbitrary_chunks( + context: Context, expects: list[tuple[int, int, int]] ) -> None: - ch1: Channel[TableChunk] = context.create_channel() + ch1: Channel[ArbitraryChunk[tuple[int, int, int]]] = context.create_channel() - # The actor access `ch1` both through the `ch_out` parameter and the closure. + # The actor accesses `ch1` both through the `ch_out` parameter and the closure. @define_actor(extra_channels=(ch1,)) async def actor1(ctx: Context, /, ch_out: Channel) -> None: for seq, chunk in enumerate(expects): - await ch1.send( - context, - Message( - seq, - TableChunk.from_pylibcudf_table( - table=chunk, - stream=stream, - exclusive_view=False, - br=context.br(), - ), - ), - ) - await ch_out.drain(context) + await ch1.send(ctx, Message(seq, ArbitraryChunk(chunk))) + await ch_out.drain(ctx) actor2, output = pull_from_channel(context, ch_in=ch1) @@ -77,18 +48,17 @@ async def actor1(ctx: Context, /, ch_out: Channel) -> None: results = output.release() for seq, (result, expect) in enumerate(zip(results, expects, strict=True)): assert result.sequence_number == seq - tbl = TableChunk.from_message(result, br=context.br()) - assert_eq(tbl.table_view(), expect) + assert ArbitraryChunk.from_message(result).release() == expect def test_shutdown(context: Context) -> None: @define_actor() - async def actor1(ctx: Context, ch_out: Channel[TableChunk]) -> None: + async def actor1(ctx: Context, ch_out: Channel[ArbitraryChunk[int]]) -> None: await ch_out.shutdown(ctx) # Calling shutdown multiple times is allowed. await ch_out.shutdown(ctx) - ch1: Channel[TableChunk] = context.create_channel() + ch1: Channel[ArbitraryChunk[int]] = context.create_channel() actor2, output = pull_from_channel(context, ch_in=ch1) run_actor_network( @@ -104,10 +74,10 @@ async def actor1(ctx: Context, ch_out: Channel[TableChunk]) -> None: def test_send_error(context: Context) -> None: @define_actor() - async def actor1(ctx: Context, ch_out: Channel[TableChunk]) -> None: + async def actor1(ctx: Context, ch_out: Channel[ArbitraryChunk[int]]) -> None: raise RuntimeError("MyError") - ch1: Channel[TableChunk] = context.create_channel() + ch1: Channel[ArbitraryChunk[int]] = context.create_channel() actor2, output = pull_from_channel(context, ch_in=ch1) with pytest.RaisesGroup( @@ -127,43 +97,38 @@ async def actor1(ctx: Context, ch_out: Channel[TableChunk]) -> None: assert output.release() == [] -def test_recv_table_chunks( - context: Context, stream: Stream, expects: list[plc.Table] +def test_recv_arbitrary_chunks( + context: Context, expects: list[tuple[int, int, int]] ) -> None: - table_chunks = [ - Message( - seq, - TableChunk.from_pylibcudf_table( - expect, stream, exclusive_view=False, br=context.br() - ), - ) - for seq, expect in enumerate(expects) + chunks = [ + Message(seq, ArbitraryChunk(expect)) for seq, expect in enumerate(expects) ] - results: list[Message[TableChunk]] = [] + results: list[Message[ArbitraryChunk[tuple[int, int, int]]]] = [] @define_actor() - async def actor1(ctx: Context, ch_in: Channel[TableChunk]) -> None: + async def actor1( + ctx: Context, ch_in: Channel[ArbitraryChunk[tuple[int, int, int]]] + ) -> None: while True: - chunk = await ch_in.recv(context) + chunk = await ch_in.recv(ctx) if chunk is None: break results.append(chunk) - ch1: Channel[TableChunk] = context.create_channel() + ch1: Channel[ArbitraryChunk[tuple[int, int, int]]] = context.create_channel() run_actor_network( context, actors=[ - push_to_channel(context, ch_out=ch1, messages=table_chunks), + push_to_channel(context, ch_out=ch1, messages=chunks), actor1(context, ch_in=ch1), ], ) for seq, (result, expect) in enumerate(zip(results, expects, strict=True)): assert result.sequence_number == seq - tbl = TableChunk.from_message(result, br=context.br()) - assert_eq(tbl.table_view(), expect) + assert ArbitraryChunk.from_message(result).release() == expect @pytest.mark.filterwarnings("error") diff --git a/python/rapidsmpf/rapidsmpf/tests/streaming/test_fanout.py b/python/rapidsmpf/rapidsmpf/tests/streaming/test_fanout.py index c41b641f8..8253589be 100644 --- a/python/rapidsmpf/rapidsmpf/tests/streaming/test_fanout.py +++ b/python/rapidsmpf/rapidsmpf/tests/streaming/test_fanout.py @@ -7,101 +7,85 @@ from typing import TYPE_CHECKING -import pylibcudf as plc import pytest -pytest.importorskip("cudf_streaming") -from cudf_streaming.streaming.table_chunk import TableChunk - +from rapidsmpf.streaming.chunks.packed_data import PackedDataChunk from rapidsmpf.streaming.core.actor import run_actor_network from rapidsmpf.streaming.core.fanout import FanoutPolicy, fanout from rapidsmpf.streaming.core.leaf_actor import pull_from_channel, push_to_channel from rapidsmpf.streaming.core.message import Message -from rapidsmpf.testing import assert_eq - -cudf = pytest.importorskip("cudf") +from rapidsmpf.testing import generate_packed_data, validate_packed_data -_INT64 = plc.DataType(plc.TypeId.INT64) +if TYPE_CHECKING: + from rapidsmpf.streaming.core.channel import Channel + from rapidsmpf.streaming.core.context import Context -def _ab_table(i: int) -> plc.Table: - return plc.Table( - [ - plc.Column.from_iterable_of_py( - [i, i + 1, i + 2], plc.DataType(plc.TypeId.INT64) - ), - plc.Column.from_iterable_of_py( - [i * 10, i * 10 + 1, i * 10 + 2], plc.DataType(plc.TypeId.INT64) - ), - ] +def _message( + context: Context, sequence_number: int, n_elements: int = 3 +) -> Message[PackedDataChunk]: + stream = context.get_stream_from_pool() + chunk = PackedDataChunk.from_packed_data( + generate_packed_data( + n_elements, + sequence_number * 10, + stream, + context.br(), + ), + br=context.br(), ) + return Message(sequence_number, chunk) -if TYPE_CHECKING: - from rmm.pylibrmm.stream import Stream - - from rapidsmpf.streaming.core.channel import Channel - from rapidsmpf.streaming.core.context import Context +def _validate( + context: Context, + msg: Message[PackedDataChunk], + sequence_number: int, + n_elements: int = 3, +) -> None: + assert msg.sequence_number == sequence_number + packed = PackedDataChunk.from_message(msg, br=context.br()).to_packed_data() + validate_packed_data(packed, n_elements, sequence_number * 10) @pytest.mark.parametrize("policy", [FanoutPolicy.BOUNDED, FanoutPolicy.UNBOUNDED]) -def test_fanout_basic(context: Context, stream: Stream, policy: FanoutPolicy) -> None: +def test_fanout_basic(context: Context, policy: FanoutPolicy) -> None: """Test basic fanout functionality with multiple output channels.""" - # Create channels - ch_in: Channel[TableChunk] = context.create_channel() - ch_out1: Channel[TableChunk] = context.create_channel() - ch_out2: Channel[TableChunk] = context.create_channel() + ch_in: Channel[PackedDataChunk] = context.create_channel() + ch_out1: Channel[PackedDataChunk] = context.create_channel() + ch_out2: Channel[PackedDataChunk] = context.create_channel() - # Create test messages - messages = [] - for i in range(5): - chunk = TableChunk.from_pylibcudf_table( - _ab_table(i), stream, exclusive_view=False, br=context.br() - ) - messages.append(Message(i, chunk)) + messages = [_message(context, i) for i in range(5)] - # Create actors push_actor = push_to_channel(context, ch_in, messages) fanout_actor = fanout(context, ch_in, [ch_out1, ch_out2], policy) pull_actor1, output1 = pull_from_channel(context, ch_out1) pull_actor2, output2 = pull_from_channel(context, ch_out2) - # Run pipeline run_actor_network( context, actors=[push_actor, fanout_actor, pull_actor1, pull_actor2], ) - # Verify results results1 = output1.release() results2 = output2.release() assert len(results1) == 5, f"Expected 5 messages in output1, got {len(results1)}" assert len(results2) == 5, f"Expected 5 messages in output2, got {len(results2)}" - # Check that both outputs received the same sequence numbers and data for i in range(5): - assert results1[i].sequence_number == i - assert results2[i].sequence_number == i - - chunk1 = TableChunk.from_message(results1[i], br=context.br()) - chunk2 = TableChunk.from_message(results2[i], br=context.br()) - - # Verify data is correct - expected_table = _ab_table(i) - assert_eq(chunk1.table_view(), expected_table) - assert_eq(chunk2.table_view(), expected_table) + _validate(context, results1[i], i) + _validate(context, results2[i], i) @pytest.mark.parametrize("num_outputs", [1, 3, 5]) @pytest.mark.parametrize("policy", [FanoutPolicy.BOUNDED, FanoutPolicy.UNBOUNDED]) def test_fanout_multiple_outputs( - context: Context, stream: Stream, num_outputs: int, policy: FanoutPolicy + context: Context, num_outputs: int, policy: FanoutPolicy ) -> None: """Test fanout with varying numbers of output channels.""" - # Create channels - ch_in: Channel[TableChunk] = context.create_channel() - chs_out: list[Channel[TableChunk]] = [ + ch_in: Channel[PackedDataChunk] = context.create_channel() + chs_out: list[Channel[PackedDataChunk]] = [ context.create_channel() for _ in range(num_outputs) ] @@ -110,22 +94,8 @@ def test_fanout_multiple_outputs( fanout(context, ch_in, chs_out, policy) return - # Create test messages - messages = [] - for i in range(3): - table = plc.Table( - [ - plc.Column.from_iterable_of_py( - [i * 10, i * 10 + 1], plc.DataType(plc.TypeId.INT64) - ), - ] - ) - chunk = TableChunk.from_pylibcudf_table( - table, stream, exclusive_view=False, br=context.br() - ) - messages.append(Message(i, chunk)) + messages = [_message(context, i, n_elements=2) for i in range(3)] - # Create actors push_actor = push_to_channel(context, ch_in, messages) fanout_actor = fanout(context, ch_in, chs_out, policy) pull_actors = [] @@ -135,25 +105,23 @@ def test_fanout_multiple_outputs( pull_actors.append(pull_actor) outputs.append(output) - # Run pipeline run_actor_network( context, actors=[push_actor, fanout_actor, *pull_actors], ) - # Verify all outputs received the messages for output_idx, output in enumerate(outputs): results = output.release() assert len(results) == 3, ( f"Output {output_idx}: Expected 3 messages, got {len(results)}" ) for i in range(3): - assert results[i].sequence_number == i + _validate(context, results[i], i, n_elements=2) -def test_fanout_empty_outputs(context: Context, stream: Stream) -> None: +def test_fanout_empty_outputs(context: Context) -> None: """Test fanout with empty output list raises value error.""" - ch_in: Channel[TableChunk] = context.create_channel() + ch_in: Channel[PackedDataChunk] = context.create_channel() with pytest.raises(ValueError): fanout(context, ch_in, [], FanoutPolicy.BOUNDED) diff --git a/python/rapidsmpf/rapidsmpf/tests/streaming/test_leaf_actor.py b/python/rapidsmpf/rapidsmpf/tests/streaming/test_leaf_actor.py index af5d36ee7..86d7fbb75 100644 --- a/python/rapidsmpf/rapidsmpf/tests/streaming/test_leaf_actor.py +++ b/python/rapidsmpf/rapidsmpf/tests/streaming/test_leaf_actor.py @@ -5,53 +5,27 @@ from typing import TYPE_CHECKING -import pylibcudf as plc -import pytest - -pytest.importorskip("cudf_streaming") -from cudf_streaming.streaming.table_chunk import TableChunk - +from rapidsmpf.streaming.chunks.arbitrary import ArbitraryChunk from rapidsmpf.streaming.core.actor import run_actor_network from rapidsmpf.streaming.core.leaf_actor import pull_from_channel, push_to_channel from rapidsmpf.streaming.core.message import Message -from rapidsmpf.testing import assert_eq - -cudf = pytest.importorskip("cudf") if TYPE_CHECKING: - from rmm.pylibrmm.stream import Stream - from rapidsmpf.streaming.core.channel import Channel from rapidsmpf.streaming.core.context import Context -def test_roundtrip(context: Context, stream: Stream) -> None: - expects = [ - plc.Table( - [ - plc.Column.from_iterable_of_py( - [1 * seq, 2 * seq, 3 * seq], plc.DataType(plc.TypeId.INT64) - ) - ] - ) - for seq in range(10) - ] - table_chunks = [ - Message( - seq, - TableChunk.from_pylibcudf_table( - expect, stream, exclusive_view=False, br=context.br() - ), - ) - for seq, expect in enumerate(expects) +def test_roundtrip(context: Context) -> None: + expects = [(seq, seq * 2, seq * 3) for seq in range(10)] + chunks = [ + Message(seq, ArbitraryChunk(expect)) for seq, expect in enumerate(expects) ] - ch1: Channel[TableChunk] = context.create_channel() - actor1 = push_to_channel(context, ch_out=ch1, messages=table_chunks) + ch1: Channel[ArbitraryChunk[tuple[int, int, int]]] = context.create_channel() + actor1 = push_to_channel(context, ch_out=ch1, messages=chunks) actor2, output = pull_from_channel(context, ch_in=ch1) run_actor_network(context, actors=(actor1, actor2)) results = output.release() for seq, (result, expect) in enumerate(zip(results, expects, strict=True)): assert result.sequence_number == seq - tbl = TableChunk.from_message(result, br=context.br()) - assert_eq(tbl.table_view(), expect) + assert ArbitraryChunk.from_message(result).release() == expect diff --git a/python/rapidsmpf/rapidsmpf/tests/streaming/test_sparse_alltoall.py b/python/rapidsmpf/rapidsmpf/tests/streaming/test_sparse_alltoall.py index db2bbe566..b1b6eb68c 100644 --- a/python/rapidsmpf/rapidsmpf/tests/streaming/test_sparse_alltoall.py +++ b/python/rapidsmpf/rapidsmpf/tests/streaming/test_sparse_alltoall.py @@ -6,37 +6,20 @@ import asyncio from typing import TYPE_CHECKING -import numpy as np -import pylibcudf as plc import pytest -pytest.importorskip("cudf_streaming") -from cudf_streaming.integrations.partition import unpack_and_concat - -from rapidsmpf.memory.packed_data import PackedData from rapidsmpf.streaming.coll.sparse_alltoall import SparseAlltoall -from rapidsmpf.testing import assert_eq - -cudf = pytest.importorskip("cudf") +from rapidsmpf.testing import generate_packed_data, validate_packed_data if TYPE_CHECKING: from rapidsmpf.communicator.communicator import Communicator + from rapidsmpf.memory.packed_data import PackedData from rapidsmpf.streaming.core.context import Context -def make_packed_data(context: Context, values: np.ndarray) -> PackedData: - stream = context.get_stream_from_pool() - table = plc.Table([plc.Column.from_array(values, stream=stream)]) - return PackedData.from_cudf_packed_columns( # type: ignore[attr-defined, no-any-return] - plc.contiguous_split.pack(table, stream=stream), - stream, - context.br(), - ) - - -def unpack_table(context: Context, packed_data: PackedData) -> plc.Table: +def make_packed_data(context: Context, value: int) -> PackedData: stream = context.get_stream_from_pool() - return unpack_and_concat([packed_data], stream, context.br()) + return generate_packed_data(1, value, stream, context.br()) def test_sparse_alltoall_non_participating_ranks( @@ -64,24 +47,13 @@ def test_sparse_alltoall_non_participating_ranks( ) if comm.rank == 0: - exchange.insert(1, make_packed_data(context, np.array([11], dtype=np.int32))) - exchange.insert(1, make_packed_data(context, np.array([29], dtype=np.int32))) + exchange.insert(1, make_packed_data(context, 11)) + exchange.insert(1, make_packed_data(context, 29)) asyncio.run(exchange.insert_finished(context)) if comm.rank == 1: results = exchange.extract(0) assert len(results) == 2 - stream = context.get_stream_from_pool() - assert_eq( - unpack_table(context, results[0]), - plc.Table( - [plc.Column.from_array(np.array([11], dtype=np.int32), stream=stream)] - ), - ) - assert_eq( - unpack_table(context, results[1]), - plc.Table( - [plc.Column.from_array(np.array([29], dtype=np.int32), stream=stream)] - ), - ) + validate_packed_data(results[0], 1, 11) + validate_packed_data(results[1], 1, 29)