From 6746a712d5157a4d69a318348b0c8f388273b02d Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Wed, 8 Jul 2026 12:23:26 -0700 Subject: [PATCH 1/8] test cases for concat vectors, extend vectors, generate vectors --- .../src/utils/containers/concat_vectors.cc | 38 +++++++++++++++++++ .../src/utils/containers/extend_vector.cc | 30 +++++++++++++++ .../src/utils/containers/generate_vector.cc | 31 +++++++++++++++ 3 files changed, 99 insertions(+) create mode 100644 lib/utils/test/src/utils/containers/concat_vectors.cc create mode 100644 lib/utils/test/src/utils/containers/extend_vector.cc create mode 100644 lib/utils/test/src/utils/containers/generate_vector.cc diff --git a/lib/utils/test/src/utils/containers/concat_vectors.cc b/lib/utils/test/src/utils/containers/concat_vectors.cc new file mode 100644 index 0000000000..e22b8e8c6c --- /dev/null +++ b/lib/utils/test/src/utils/containers/concat_vectors.cc @@ -0,0 +1,38 @@ +#include "utils/containers/concat_vectors.h" +#include "test/utils/doctest/fmt/vector.h" +#include +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("concat_vectors") { + SUBCASE("two vectors") { + std::vector prefix = {1, 2}; + std::vector postfix = {3, 4}; + + std::vector result = concat_vectors(prefix, postfix); + std::vector correct = {1, 2, 3, 4}; + + CHECK(result == correct); + } + + SUBCASE("vector of vectors") { + std::vector> vecs = {{1, 2}, {3}, {4, 5}}; + + std::vector result = concat_vectors(vecs); + std::vector correct = {1, 2, 3, 4, 5}; + + CHECK(result == correct); + } + + SUBCASE("empty vector of vectors") { + std::vector> vecs = {}; + + std::vector result = concat_vectors(vecs); + std::vector correct = {}; + + CHECK(result == correct); + } + } +} diff --git a/lib/utils/test/src/utils/containers/extend_vector.cc b/lib/utils/test/src/utils/containers/extend_vector.cc new file mode 100644 index 0000000000..a394cdf22f --- /dev/null +++ b/lib/utils/test/src/utils/containers/extend_vector.cc @@ -0,0 +1,30 @@ +#include "utils/containers/extend_vector.h" +#include "test/utils/doctest/fmt/vector.h" +#include +#include + +using namespace ::FlexFlow; +// checks rhs gets added to the end of lhs +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("extend_vector") { + SUBCASE("non-empty vectors") { + std::vector lhs = {1, 2}; + std::vector rhs = {3, 4}; + + extend_vector(lhs, rhs); + std::vector correct = {1, 2, 3, 4}; + + CHECK(lhs == correct); + } + // checks that lhs stays the same when the rhs is empty + SUBCASE("empty rhs") { + std::vector lhs = {1, 2}; + std::vector rhs = {}; + + extend_vector(lhs, rhs); + std::vector correct = {1, 2}; + + CHECK(lhs == correct); + } + } +} diff --git a/lib/utils/test/src/utils/containers/generate_vector.cc b/lib/utils/test/src/utils/containers/generate_vector.cc new file mode 100644 index 0000000000..e560ba80bd --- /dev/null +++ b/lib/utils/test/src/utils/containers/generate_vector.cc @@ -0,0 +1,31 @@ +#include "utils/containers/generate_vector.h" +#include "test/utils/doctest/fmt/vector.h" +#include "utils/exception.h" +#include "utils/nonnegative_int/nonnegative_int.h" +#include +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("generate_vector") { + SUBCASE("empty vector") { + std::vector result = generate_vector(0_n, [](nonnegative_int) -> int { + PANIC("lambda should not be called"); + }); + std::vector correct = {}; + + CHECK(result == correct); + } + SUBCASE("non-empty vector") { + std::vector result = generate_vector(5_n, [](nonnegative_int idx) { + int i = idx.unwrap_nonnegative(); + return i * i; + + }); + std::vector correct = {0, 1, 4, 9, 16}; + + CHECK(result == correct); + } + } +} From 0937720eaa1d55a9d640aad8e4b7fcd4022a4f6a Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Wed, 22 Jul 2026 01:28:39 -0700 Subject: [PATCH 2/8] Add NCCL hello world Realm task --- .../realm-execution/tasks/impl/nccl_task.h | 24 ++++++++ .../tasks/impl/nccl_task_args.dtg.toml | 12 ++++ .../impl/serializable_nccl_task_args.dtg.toml | 13 +++++ .../tasks/impl/serializable_nccl_task_args.h | 17 ++++++ .../realm-execution/tasks/task_id_t.dtg.toml | 3 + .../realm-execution/tasks/impl/nccl_task.cc | 57 +++++++++++++++++++ .../tasks/impl/serializable_nccl_task_args.cc | 19 +++++++ .../tasks/realm_task_registry.cc | 6 ++ .../test/src/realm-execution/nccl_task.cc | 36 ++++++++++++ 9 files changed, 187 insertions(+) create mode 100644 lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h create mode 100644 lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml create mode 100644 lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml create mode 100644 lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h create mode 100644 lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc create mode 100644 lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc create mode 100644 lib/realm-execution/test/src/realm-execution/nccl_task.cc diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h new file mode 100644 index 0000000000..cecf444eb3 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -0,0 +1,24 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H + +#include "realm-execution/realm.h" +#include "realm-execution/realm_context.h" +#include +#include + +namespace FlexFlow { + void nccl_task_body(void const *args, + size_t arglen, + void const *userdata, + size_t userdata_len, + Realm::Processor proc); + + +Realm::Event spawn_nccl_task(RealmContext &ctx, + Realm::Processor target_proc, + std::string const &message, + Realm::Event precondition); + +} + +#endif diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml new file mode 100644 index 0000000000..fa1d7a3c08 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml @@ -0,0 +1,12 @@ +namespace = "FlexFlow" +name = "NcclTaskArgs" +type = "struct" +features = [] + +includes = [ + "string", +] + +[[fields]] +name = "message" +type = "std::string" diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml new file mode 100644 index 0000000000..b77585ee7c --- /dev/null +++ b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml @@ -0,0 +1,13 @@ +namespace = "FlexFlow" +name = "SerializableNcclTaskArgs" +type = "struct" +features = [ + "json", +] +includes = [ + "string", +] + +[[fields]] +name = "message" +type = "std::string" diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h new file mode 100644 index 0000000000..9b90be11c2 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h @@ -0,0 +1,17 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_SERIALIZABLE_NCCL_TASK_ARGS_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_SERIALIZABLE_NCCL_TASK_ARGS_H + +#include "realm-execution/tasks/impl/nccl_task_args.dtg.h" +#include "realm-execution/tasks/impl/serializable_nccl_task_args.dtg.h" + +namespace FlexFlow { + +SerializableNcclTaskArgs + nccl_task_args_to_serializable(NcclTaskArgs const &); + +NcclTaskArgs + nccl_task_args_from_serializable(SerializableNcclTaskArgs const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml index a7c2d54180..10961065da 100644 --- a/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml @@ -270,6 +270,9 @@ name = "NCCL_GETUNIQUEID_TASK_ID" [[values]] name = "NCCL_INIT_COMMS_TASK_ID" +[[values]] +name = "NCCL_HELLO_WORLD_TASK_ID" + [[values]] name = "STRATEGY_SEARCH_TASK_ID" diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc new file mode 100644 index 0000000000..be6b232f46 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc @@ -0,0 +1,57 @@ +#include "realm-execution/tasks/impl/nccl_task.h" + +#include "realm-execution/tasks/impl/nccl_task_args.dtg.h" +#include "realm-execution/tasks/impl/serializable_nccl_task_args.h" +#include "realm-execution/tasks/serializer/task_arg_serializer.h" +#include "realm-execution/tasks/task_id_t.h" + +#include +#include + +namespace FlexFlow { + +void nccl_task_body(void const *args, + size_t arglen, + void const *userdata, + size_t userdata_len, + Realm::Processor proc) { + (void)userdata; + (void)userdata_len; + (void)proc; + + NcclTaskArgs task_args = nccl_task_args_from_serializable( + deserialize_task_args(args, arglen)); + + int nccl_version = 0; + ncclResult_t result = ncclGetVersion(&nccl_version); + + if (result != ncclSuccess) { + std::printf("NCCL error: %s\n", ncclGetErrorString(result)); + return; + } + + std::printf("%s\n", task_args.message.c_str()); + std::printf("NCCL version: %d\n", nccl_version); +} + +Realm::Event spawn_nccl_task(RealmContext &ctx, + Realm::Processor target_proc, + std::string const &message, + Realm::Event precondition) { + NcclTaskArgs task_args = NcclTaskArgs{ + /*message=*/message, + }; + + std::string serialized_args = + serialize_task_args(nccl_task_args_to_serializable(task_args)); + + return ctx.spawn_task( + target_proc, + task_id_t::NCCL_HELLO_WORLD_TASK_ID, + serialized_args.data(), + serialized_args.size(), + Realm::ProfilingRequestSet{}, + precondition); +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc new file mode 100644 index 0000000000..57f6cd1c1b --- /dev/null +++ b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc @@ -0,0 +1,19 @@ +#include "realm-execution/tasks/impl/serializable_nccl_task_args.h" + +namespace FlexFlow { + +SerializableNcclTaskArgs + nccl_task_args_to_serializable(NcclTaskArgs const &args) { + return SerializableNcclTaskArgs{ + args.message, + }; +} + +NcclTaskArgs + nccl_task_args_from_serializable(SerializableNcclTaskArgs const &args) { + return NcclTaskArgs{ + args.message, + }; +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc index 406a2912a3..0fa6769eca 100644 --- a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc +++ b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc @@ -7,6 +7,7 @@ #include "realm-execution/tasks/impl/per_device_op_state_init_task.h" #include "realm-execution/tasks/task_id_t.h" #include "utils/exception.h" +#include "realm-execution/tasks/impl/nccl_task.h" namespace FlexFlow { @@ -133,6 +134,11 @@ Realm::Event register_all_tasks() { register_task(Realm::Processor::TOC_PROC, task_id, op_task_body)); } + pending_registrations.push_back( + register_task(Realm::Processor::TOC_PROC, + task_id_t::NCCL_HELLO_WORLD_TASK_ID, + nccl_task_body)); + pending_registrations.push_back(register_task(Realm::Processor::LOC_PROC, task_id_t::CONTROLLER_TASK_ID, controller_task_body)); diff --git a/lib/realm-execution/test/src/realm-execution/nccl_task.cc b/lib/realm-execution/test/src/realm-execution/nccl_task.cc new file mode 100644 index 0000000000..79a34fdac1 --- /dev/null +++ b/lib/realm-execution/test/src/realm-execution/nccl_task.cc @@ -0,0 +1,36 @@ +#include "internal/realm_test_utils.h" +#include "realm-execution/realm_manager.h" +#include "realm-execution/tasks/impl/nccl_task.h" +#include + +namespace test { + +using namespace ::FlexFlow; +namespace Realm = ::FlexFlow::Realm; + +TEST_SUITE(FF_CUDA_TEST_SUITE) { + TEST_CASE("NCCL task prints Hello World") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/1_n); + + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager(&fake_argc, &fake_argv); + + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + Realm::Event event = spawn_nccl_task( + ctx, + ctx.get_current_processor(), + "Hello World from NCCL!", + Realm::Event::NO_EVENT); + + event.wait(); + }); + + result.wait(); + } +} + +} // namespace test From 598822f8bd3f663ec8e1ee08ee049002b714e343 Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Tue, 28 Jul 2026 11:05:17 -0700 Subject: [PATCH 3/8] Add helper for NCCL AllReduce --- .../realm-execution/tasks/impl/nccl_task.h | 20 ++++++++++++++----- .../realm-execution/tasks/impl/nccl_task.cc | 18 ++++++++++++++++- 2 files changed, 32 insertions(+), 6 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h index cecf444eb3..ec875c7901 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -1,17 +1,27 @@ #ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H #define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H +#include "kernels/device.h" +#include #include "realm-execution/realm.h" #include "realm-execution/realm_context.h" #include #include namespace FlexFlow { - void nccl_task_body(void const *args, - size_t arglen, - void const *userdata, - size_t userdata_len, - Realm::Processor proc); + +ncclResult_t run_nccl_all_reduce(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + ncclComm_t communicator, + ffStream_t stream); +void nccl_task_body(void const *args, + size_t arglen, + void const *userdata, + size_t userdata_len, + Realm::Processor proc); Realm::Event spawn_nccl_task(RealmContext &ctx, diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc index be6b232f46..57f1796709 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc @@ -1,5 +1,5 @@ #include "realm-execution/tasks/impl/nccl_task.h" - +#include "kernels/device.h" #include "realm-execution/tasks/impl/nccl_task_args.dtg.h" #include "realm-execution/tasks/impl/serializable_nccl_task_args.h" #include "realm-execution/tasks/serializer/task_arg_serializer.h" @@ -10,6 +10,22 @@ namespace FlexFlow { +ncclResult_t run_nccl_all_reduce(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + ncclComm_t communicator, + ffStream_t stream) { + return ncclAllReduce(send_buffer, + receive_buffer, + count, + data_type, + reduction_op, + communicator, + stream); +} + void nccl_task_body(void const *args, size_t arglen, void const *userdata, From f2a4577412749ea9fb3532904c9e2272d18abdba Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Fri, 31 Jul 2026 01:55:21 -0700 Subject: [PATCH 4/8] Add NCCL broadcast and reduce helpers --- .../realm-execution/tasks/impl/nccl_task.h | 18 +++++++++- .../realm-execution/tasks/impl/nccl_task.cc | 35 ++++++++++++++++++- 2 files changed, 51 insertions(+), 2 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h index ec875c7901..9e3601520d 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -17,13 +17,29 @@ ncclResult_t run_nccl_all_reduce(void const *send_buffer, ncclRedOp_t reduction_op, ncclComm_t communicator, ffStream_t stream); + +ncclResult_t run_nccl_broadcast(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + ncclComm_t communicator, + ffStream_t stream); + +ncclResult_t run_nccl_reduce(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + ncclComm_t communicator, + ffStream_t stream); + void nccl_task_body(void const *args, size_t arglen, void const *userdata, size_t userdata_len, Realm::Processor proc); - Realm::Event spawn_nccl_task(RealmContext &ctx, Realm::Processor target_proc, std::string const &message, diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc index 57f1796709..070ac3c937 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc @@ -1,5 +1,4 @@ #include "realm-execution/tasks/impl/nccl_task.h" -#include "kernels/device.h" #include "realm-execution/tasks/impl/nccl_task_args.dtg.h" #include "realm-execution/tasks/impl/serializable_nccl_task_args.h" #include "realm-execution/tasks/serializer/task_arg_serializer.h" @@ -26,6 +25,40 @@ ncclResult_t run_nccl_all_reduce(void const *send_buffer, stream); } +ncclResult_t run_nccl_broadcast(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + int root_rank, + ncclComm_t communicator, + ffStream_t stream) { + return ncclBroadcast(send_buffer, + receive_buffer, + count, + data_type, + root_rank, + communicator, + stream); +} + +ncclResult_t run_nccl_reduce(void const *send_buffer, + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + int root_rank, + ncclComm_t communicator, + ffStream_t stream) { + return ncclReduce(send_buffer, + receive_buffer, + count, + data_type, + reduction_op, + root_rank, + communicator, + stream); +} + void nccl_task_body(void const *args, size_t arglen, void const *userdata, From cfede62ade4ddf227c5c86e800133e868b2d9db3 Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Wed, 5 Aug 2026 01:17:48 -0700 Subject: [PATCH 5/8] Add NCCL broadcast and reduce tests --- .../realm-execution/tasks/impl/nccl_task.h | 15 ++-- .../test/src/realm-execution/nccl_task.cc | 88 +++++++++++++++++++ 2 files changed, 96 insertions(+), 7 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h index 9e3601520d..506ea43efd 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -22,17 +22,18 @@ ncclResult_t run_nccl_broadcast(void const *send_buffer, void *receive_buffer, size_t count, ncclDataType_t data_type, - ncclRedOp_t reduction_op, + int root_rank, ncclComm_t communicator, ffStream_t stream); ncclResult_t run_nccl_reduce(void const *send_buffer, - void *receive_buffer, - size_t count, - ncclDataType_t data_type, - ncclRedOp_t reduction_op, - ncclComm_t communicator, - ffStream_t stream); + void *receive_buffer, + size_t count, + ncclDataType_t data_type, + ncclRedOp_t reduction_op, + int root_rank, + ncclComm_t communicator, + ffStream_t stream); void nccl_task_body(void const *args, size_t arglen, diff --git a/lib/realm-execution/test/src/realm-execution/nccl_task.cc b/lib/realm-execution/test/src/realm-execution/nccl_task.cc index 79a34fdac1..9dd94ebf9a 100644 --- a/lib/realm-execution/test/src/realm-execution/nccl_task.cc +++ b/lib/realm-execution/test/src/realm-execution/nccl_task.cc @@ -1,7 +1,11 @@ #include "internal/realm_test_utils.h" #include "realm-execution/realm_manager.h" #include "realm-execution/tasks/impl/nccl_task.h" + +#include #include +#include +#include namespace test { @@ -31,6 +35,90 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { result.wait(); } + + TEST_CASE("NCCL broadcast and reduce helpers") { + constexpr size_t count = 8; + size_t const buffer_size = count * sizeof(int); + + ncclUniqueId unique_id; + REQUIRE(ncclGetUniqueId(&unique_id) == ncclSuccess); + + ncclComm_t communicator; + REQUIRE(ncclCommInitRank( + &communicator, + /*num_ranks=*/1, + unique_id, + /*rank=*/0) == ncclSuccess); + + ffStream_t stream; + REQUIRE(cudaStreamCreate(&stream) == cudaSuccess); + + int *send_buffer = nullptr; + int *receive_buffer = nullptr; + + REQUIRE(cudaMalloc(&send_buffer, buffer_size) == cudaSuccess); + REQUIRE(cudaMalloc(&receive_buffer, buffer_size) == cudaSuccess); + + std::vector input = {1, 2, 3, 4, 5, 6, 7, 8}; + std::vector output(count, 0); + + REQUIRE(cudaMemcpy(send_buffer, + input.data(), + buffer_size, + cudaMemcpyHostToDevice) == cudaSuccess); + + SUBCASE("broadcast") { + REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); + + REQUIRE(run_nccl_broadcast(send_buffer, + receive_buffer, + count, + ncclInt32, + /*root_rank=*/0, + communicator, + stream) == ncclSuccess); + + REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); + + REQUIRE(cudaMemcpy(output.data(), + receive_buffer, + buffer_size, + cudaMemcpyDeviceToHost) == cudaSuccess); + + for (size_t i = 0; i < count; i++) { + CHECK(output[i] == input[i]); + } + } + + SUBCASE("reduce") { + REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); + + REQUIRE(run_nccl_reduce(send_buffer, + receive_buffer, + count, + ncclInt32, + ncclSum, + /*root_rank=*/0, + communicator, + stream) == ncclSuccess); + + REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); + + REQUIRE(cudaMemcpy(output.data(), + receive_buffer, + buffer_size, + cudaMemcpyDeviceToHost) == cudaSuccess); + + for (size_t i = 0; i < count; i++) { + CHECK(output[i] == input[i]); + } + } + + REQUIRE(cudaFree(send_buffer) == cudaSuccess); + REQUIRE(cudaFree(receive_buffer) == cudaSuccess); + REQUIRE(cudaStreamDestroy(stream) == cudaSuccess); + REQUIRE(ncclCommDestroy(communicator) == ncclSuccess); + } } } // namespace test From 52536e1e39f56297ce3f81ad7944beb82cc076df Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Fri, 7 Aug 2026 19:05:10 -0700 Subject: [PATCH 6/8] Add spawn-level NCCL task test with real task args --- .../realm-execution/tasks/impl/nccl_task.h | 8 +- .../tasks/impl/nccl_task_args.dtg.toml | 11 +- .../impl/serializable_nccl_task_args.dtg.toml | 12 +- .../realm-execution/tasks/impl/nccl_task.cc | 14 +- .../tasks/impl/serializable_nccl_task_args.cc | 16 +- .../tasks/realm_task_registry.cc | 5 + .../test/src/realm-execution/nccl_task.cc | 193 ++++++++++-------- 7 files changed, 157 insertions(+), 102 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h index 506ea43efd..ae5c017b61 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -1,6 +1,8 @@ #ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H #define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H +#include "realm-execution/tensor_instance_backing.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "kernels/device.h" #include #include "realm-execution/realm.h" @@ -41,9 +43,11 @@ void nccl_task_body(void const *args, size_t userdata_len, Realm::Processor proc); -Realm::Event spawn_nccl_task(RealmContext &ctx, +Realm::Event spawn_nccl_task( + RealmContext &ctx, Realm::Processor target_proc, - std::string const &message, + DynamicNodeInvocation const &invocation, + TensorInstanceBacking const &tensor_backing, Realm::Event precondition); } diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml index fa1d7a3c08..304db124ff 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml @@ -4,9 +4,14 @@ type = "struct" features = [] includes = [ - "string", + "realm-execution/tensor_instance_backing.dtg.h", + "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", ] [[fields]] -name = "message" -type = "std::string" +name = "invocation" +type = "::FlexFlow::DynamicNodeInvocation" + +[[fields]] +name = "tensor_backing" +type = "::FlexFlow::TensorInstanceBacking" diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml index b77585ee7c..ebdc97cbca 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml @@ -4,10 +4,16 @@ type = "struct" features = [ "json", ] + includes = [ - "string", + "realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.h", + "task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.h", ] [[fields]] -name = "message" -type = "std::string" +name = "invocation" +type = "::FlexFlow::SerializableDynamicNodeInvocation" + +[[fields]] +name = "tensor_backing" +type = "::FlexFlow::SerializableTensorInstanceBacking" diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc index 070ac3c937..49d4c321fc 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc @@ -79,16 +79,18 @@ void nccl_task_body(void const *args, return; } - std::printf("%s\n", task_args.message.c_str()); std::printf("NCCL version: %d\n", nccl_version); } -Realm::Event spawn_nccl_task(RealmContext &ctx, - Realm::Processor target_proc, - std::string const &message, - Realm::Event precondition) { +Realm::Event spawn_nccl_task( + RealmContext &ctx, + Realm::Processor target_proc, + DynamicNodeInvocation const &invocation, + TensorInstanceBacking const &tensor_backing, + Realm::Event precondition) { NcclTaskArgs task_args = NcclTaskArgs{ - /*message=*/message, + /*invocation=*/invocation, + /*tensor_backing=*/tensor_backing, }; std::string serialized_args = diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc index 57f6cd1c1b..c00d56f07c 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc @@ -1,18 +1,26 @@ #include "realm-execution/tasks/impl/serializable_nccl_task_args.h" +#include "realm-execution/tasks/serializer/serializable_tensor_instance_backing.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" namespace FlexFlow { SerializableNcclTaskArgs - nccl_task_args_to_serializable(NcclTaskArgs const &args) { +nccl_task_args_to_serializable(NcclTaskArgs const &args) { return SerializableNcclTaskArgs{ - args.message, + /*invocation=*/ + dynamic_node_invocation_to_serializable(args.invocation), + /*tensor_backing=*/ + tensor_instance_backing_to_serializable(args.tensor_backing), }; } NcclTaskArgs - nccl_task_args_from_serializable(SerializableNcclTaskArgs const &args) { +nccl_task_args_from_serializable(SerializableNcclTaskArgs const &args) { return NcclTaskArgs{ - args.message, + /*invocation=*/ + dynamic_node_invocation_from_serializable(args.invocation), + /*tensor_backing=*/ + tensor_instance_backing_from_serializable(args.tensor_backing), }; } diff --git a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc index 0fa6769eca..1f79036d5c 100644 --- a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc +++ b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc @@ -134,6 +134,11 @@ Realm::Event register_all_tasks() { register_task(Realm::Processor::TOC_PROC, task_id, op_task_body)); } + pending_registrations.push_back( + register_task(Realm::Processor::LOC_PROC, + task_id_t::NCCL_HELLO_WORLD_TASK_ID, + nccl_task_body)); + pending_registrations.push_back( register_task(Realm::Processor::TOC_PROC, task_id_t::NCCL_HELLO_WORLD_TASK_ID, diff --git a/lib/realm-execution/test/src/realm-execution/nccl_task.cc b/lib/realm-execution/test/src/realm-execution/nccl_task.cc index 9dd94ebf9a..eb3add6769 100644 --- a/lib/realm-execution/test/src/realm-execution/nccl_task.cc +++ b/lib/realm-execution/test/src/realm-execution/nccl_task.cc @@ -1,11 +1,11 @@ #include "internal/realm_test_utils.h" #include "realm-execution/realm_manager.h" #include "realm-execution/tasks/impl/nccl_task.h" +#include "realm-execution/tensor_instance_backing.h" #include #include #include -#include namespace test { @@ -13,112 +13,137 @@ using namespace ::FlexFlow; namespace Realm = ::FlexFlow::Realm; TEST_SUITE(FF_CUDA_TEST_SUITE) { - TEST_CASE("NCCL task prints Hello World") { - std::vector fake_args = - make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/1_n); - int fake_argc = fake_args.size(); - char **fake_argv = fake_args.data(); +TEST_CASE("NCCL task spawns successfully") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/0_n); + + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager(&fake_argc, &fake_argv); + + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + DynamicNodeInvocation invocation{ + /*inputs=*/{}, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/std::nullopt, + /*layer_guid=*/ + dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{0}, + }, + }, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{}, + }; + + TensorInstanceBacking tensor_backing = + make_empty_tensor_instance_backing(); + + Realm::Event event = spawn_nccl_task( + ctx, + ctx.get_current_processor(), + invocation, + tensor_backing, + Realm::Event::NO_EVENT); + + event.wait(); + }); + + result.wait(); +} - RealmManager manager(&fake_argc, &fake_argv); +TEST_CASE("NCCL broadcast and reduce helpers") { + constexpr size_t count = 8; + size_t const buffer_size = count * sizeof(int); - ControllerTaskResult result = - manager.start_controller([](RealmContext &ctx) { - Realm::Event event = spawn_nccl_task( - ctx, - ctx.get_current_processor(), - "Hello World from NCCL!", - Realm::Event::NO_EVENT); + ncclUniqueId unique_id; + REQUIRE(ncclGetUniqueId(&unique_id) == ncclSuccess); - event.wait(); - }); + ncclComm_t communicator; + REQUIRE(ncclCommInitRank( + &communicator, + /*num_ranks=*/1, + unique_id, + /*rank=*/0) == ncclSuccess); - result.wait(); - } + ffStream_t stream; + REQUIRE(cudaStreamCreate(&stream) == cudaSuccess); - TEST_CASE("NCCL broadcast and reduce helpers") { - constexpr size_t count = 8; - size_t const buffer_size = count * sizeof(int); + int *send_buffer = nullptr; + int *receive_buffer = nullptr; - ncclUniqueId unique_id; - REQUIRE(ncclGetUniqueId(&unique_id) == ncclSuccess); + REQUIRE(cudaMalloc(&send_buffer, buffer_size) == cudaSuccess); + REQUIRE(cudaMalloc(&receive_buffer, buffer_size) == cudaSuccess); - ncclComm_t communicator; - REQUIRE(ncclCommInitRank( - &communicator, - /*num_ranks=*/1, - unique_id, - /*rank=*/0) == ncclSuccess); + std::vector input = {1, 2, 3, 4, 5, 6, 7, 8}; + std::vector output(count, 0); - ffStream_t stream; - REQUIRE(cudaStreamCreate(&stream) == cudaSuccess); + REQUIRE(cudaMemcpy(send_buffer, + input.data(), + buffer_size, + cudaMemcpyHostToDevice) == cudaSuccess); - int *send_buffer = nullptr; - int *receive_buffer = nullptr; + SUBCASE("broadcast") { + REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); - REQUIRE(cudaMalloc(&send_buffer, buffer_size) == cudaSuccess); - REQUIRE(cudaMalloc(&receive_buffer, buffer_size) == cudaSuccess); + REQUIRE(run_nccl_broadcast(send_buffer, + receive_buffer, + count, + ncclInt32, + /*root_rank=*/0, + communicator, + stream) == ncclSuccess); - std::vector input = {1, 2, 3, 4, 5, 6, 7, 8}; - std::vector output(count, 0); + REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); - REQUIRE(cudaMemcpy(send_buffer, - input.data(), + REQUIRE(cudaMemcpy(output.data(), + receive_buffer, buffer_size, - cudaMemcpyHostToDevice) == cudaSuccess); - - SUBCASE("broadcast") { - REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); - - REQUIRE(run_nccl_broadcast(send_buffer, - receive_buffer, - count, - ncclInt32, - /*root_rank=*/0, - communicator, - stream) == ncclSuccess); - - REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); - - REQUIRE(cudaMemcpy(output.data(), - receive_buffer, - buffer_size, - cudaMemcpyDeviceToHost) == cudaSuccess); + cudaMemcpyDeviceToHost) == cudaSuccess); - for (size_t i = 0; i < count; i++) { - CHECK(output[i] == input[i]); - } + for (size_t i = 0; i < count; i++) { + CHECK(output[i] == input[i]); } + } - SUBCASE("reduce") { - REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); + SUBCASE("reduce") { + REQUIRE(cudaMemset(receive_buffer, 0, buffer_size) == cudaSuccess); - REQUIRE(run_nccl_reduce(send_buffer, - receive_buffer, - count, - ncclInt32, - ncclSum, - /*root_rank=*/0, - communicator, - stream) == ncclSuccess); + REQUIRE(run_nccl_reduce(send_buffer, + receive_buffer, + count, + ncclInt32, + ncclSum, + /*root_rank=*/0, + communicator, + stream) == ncclSuccess); - REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); + REQUIRE(cudaStreamSynchronize(stream) == cudaSuccess); - REQUIRE(cudaMemcpy(output.data(), - receive_buffer, - buffer_size, - cudaMemcpyDeviceToHost) == cudaSuccess); + REQUIRE(cudaMemcpy(output.data(), + receive_buffer, + buffer_size, + cudaMemcpyDeviceToHost) == cudaSuccess); - for (size_t i = 0; i < count; i++) { - CHECK(output[i] == input[i]); - } + for (size_t i = 0; i < count; i++) { + CHECK(output[i] == input[i]); } - - REQUIRE(cudaFree(send_buffer) == cudaSuccess); - REQUIRE(cudaFree(receive_buffer) == cudaSuccess); - REQUIRE(cudaStreamDestroy(stream) == cudaSuccess); - REQUIRE(ncclCommDestroy(communicator) == ncclSuccess); } + + REQUIRE(cudaFree(send_buffer) == cudaSuccess); + REQUIRE(cudaFree(receive_buffer) == cudaSuccess); + REQUIRE(cudaStreamDestroy(stream) == cudaSuccess); + REQUIRE(ncclCommDestroy(communicator) == ncclSuccess); } +} // TEST_SUITE + } // namespace test From f307da5bc9d40d7f1b1de3191c9ae1233903e3fc Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Fri, 7 Aug 2026 21:41:14 -0700 Subject: [PATCH 7/8] Add broadcast attribute handling to NCCL task --- .../test/src/realm-execution/nccl_task.cc | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/lib/realm-execution/test/src/realm-execution/nccl_task.cc b/lib/realm-execution/test/src/realm-execution/nccl_task.cc index eb3add6769..602da18586 100644 --- a/lib/realm-execution/test/src/realm-execution/nccl_task.cc +++ b/lib/realm-execution/test/src/realm-execution/nccl_task.cc @@ -1,7 +1,10 @@ #include "internal/realm_test_utils.h" +#include "op-attrs/ops/broadcast_attrs.dtg.h" +#include "op-attrs/pcg_operator_attrs.dtg.h" #include "realm-execution/realm_manager.h" #include "realm-execution/tasks/impl/nccl_task.h" #include "realm-execution/tensor_instance_backing.h" +#include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" #include #include @@ -32,7 +35,18 @@ TEST_CASE("NCCL task spawns successfully") { /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/ + TrainingOperationAttrs{ + PCGOperatorAttrs{ + BroadcastAttrs{ + TensorDims{ + FFOrdered{ + 8_p, + }, + }, + }, + }, + }, /*layer_guid=*/ dynamic_layer_guid_t{ parallel_layer_guid_t{ From 542c014974c0199a23ee3b540fae125f1e300b12 Mon Sep 17 00:00:00 2001 From: Kim Hoang Date: Fri, 7 Aug 2026 23:50:21 -0700 Subject: [PATCH 8/8] Pass device handle to NCCL task --- .../realm-execution/tasks/impl/nccl_task.h | 2 ++ .../tasks/impl/nccl_task_args.dtg.toml | 8 +++++- .../impl/serializable_nccl_task_args.dtg.toml | 5 ++++ .../tasks/impl/serializable_nccl_task_args.h | 4 +-- .../realm-execution/tasks/impl/nccl_task.cc | 6 ++-- .../tasks/impl/serializable_nccl_task_args.cc | 12 ++++++-- .../test/src/realm-execution/nccl_task.cc | 28 +++++++++++++------ 7 files changed, 48 insertions(+), 17 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h index ae5c017b61..bc146a30f4 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task.h @@ -1,6 +1,7 @@ #ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H #define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_IMPL_NCCL_TASK_H +#include "realm-execution/device_specific_managed_per_device_ff_handle.h" #include "realm-execution/tensor_instance_backing.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "kernels/device.h" @@ -48,6 +49,7 @@ Realm::Event spawn_nccl_task( Realm::Processor target_proc, DynamicNodeInvocation const &invocation, TensorInstanceBacking const &tensor_backing, + DeviceSpecificPtr const &device_handle, Realm::Event precondition); } diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml index 304db124ff..b819b794df 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/impl/nccl_task_args.dtg.toml @@ -1,9 +1,11 @@ namespace = "FlexFlow" -name = "NcclTaskArgs" +name = "NCCLTaskArgs" type = "struct" features = [] includes = [ + "realm-execution/device_specific_managed_per_device_ff_handle.h", + "realm-execution/device_specific_ptr.h", "realm-execution/tensor_instance_backing.dtg.h", "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", ] @@ -15,3 +17,7 @@ type = "::FlexFlow::DynamicNodeInvocation" [[fields]] name = "tensor_backing" type = "::FlexFlow::TensorInstanceBacking" + +[[fields]] +name = "device_handle" +type = "::FlexFlow::DeviceSpecificPtr<::FlexFlow::ManagedPerDeviceFFHandle>" diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml index ebdc97cbca..919913fd79 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.dtg.toml @@ -6,6 +6,7 @@ features = [ ] includes = [ + "realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.h", "realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.h", "task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.h", ] @@ -17,3 +18,7 @@ type = "::FlexFlow::SerializableDynamicNodeInvocation" [[fields]] name = "tensor_backing" type = "::FlexFlow::SerializableTensorInstanceBacking" + +[[fields]] +name = "device_handle" +type = "::FlexFlow::SerializableDeviceSpecificPtr" diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h index 9b90be11c2..9e4eceef68 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h +++ b/lib/realm-execution/include/realm-execution/tasks/impl/serializable_nccl_task_args.h @@ -7,9 +7,9 @@ namespace FlexFlow { SerializableNcclTaskArgs - nccl_task_args_to_serializable(NcclTaskArgs const &); + nccl_task_args_to_serializable(NCCLTaskArgs const &); -NcclTaskArgs +NCCLTaskArgs nccl_task_args_from_serializable(SerializableNcclTaskArgs const &); } // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc index 49d4c321fc..05d7a982f3 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/nccl_task.cc @@ -68,7 +68,7 @@ void nccl_task_body(void const *args, (void)userdata_len; (void)proc; - NcclTaskArgs task_args = nccl_task_args_from_serializable( + NCCLTaskArgs task_args = nccl_task_args_from_serializable( deserialize_task_args(args, arglen)); int nccl_version = 0; @@ -87,10 +87,12 @@ Realm::Event spawn_nccl_task( Realm::Processor target_proc, DynamicNodeInvocation const &invocation, TensorInstanceBacking const &tensor_backing, + DeviceSpecificPtr const &device_handle, Realm::Event precondition) { - NcclTaskArgs task_args = NcclTaskArgs{ + NCCLTaskArgs task_args = NCCLTaskArgs{ /*invocation=*/invocation, /*tensor_backing=*/tensor_backing, + /*device_handle=*/device_handle, }; std::string serialized_args = diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc index c00d56f07c..c19cbc0c46 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/serializable_nccl_task_args.cc @@ -1,26 +1,32 @@ #include "realm-execution/tasks/impl/serializable_nccl_task_args.h" +#include "realm-execution/tasks/serializer/serializable_device_specific_ptr.h" #include "realm-execution/tasks/serializer/serializable_tensor_instance_backing.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" namespace FlexFlow { SerializableNcclTaskArgs -nccl_task_args_to_serializable(NcclTaskArgs const &args) { +nccl_task_args_to_serializable(NCCLTaskArgs const &args) { return SerializableNcclTaskArgs{ /*invocation=*/ dynamic_node_invocation_to_serializable(args.invocation), /*tensor_backing=*/ tensor_instance_backing_to_serializable(args.tensor_backing), + /*device_handle=*/ + device_specific_ptr_to_serializable(args.device_handle), }; } -NcclTaskArgs +NCCLTaskArgs nccl_task_args_from_serializable(SerializableNcclTaskArgs const &args) { - return NcclTaskArgs{ + return NCCLTaskArgs{ /*invocation=*/ dynamic_node_invocation_from_serializable(args.invocation), /*tensor_backing=*/ tensor_instance_backing_from_serializable(args.tensor_backing), + /*device_handle=*/ + device_specific_ptr_from_serializable( + args.device_handle), }; } diff --git a/lib/realm-execution/test/src/realm-execution/nccl_task.cc b/lib/realm-execution/test/src/realm-execution/nccl_task.cc index 602da18586..e4af71e716 100644 --- a/lib/realm-execution/test/src/realm-execution/nccl_task.cc +++ b/lib/realm-execution/test/src/realm-execution/nccl_task.cc @@ -1,5 +1,6 @@ #include "internal/realm_test_utils.h" -#include "op-attrs/ops/broadcast_attrs.dtg.h" +#include "realm-execution/distributed_ff_handle.h" +#include "op-attrs/ops/replicate_attrs.dtg.h" #include "op-attrs/pcg_operator_attrs.dtg.h" #include "realm-execution/realm_manager.h" #include "realm-execution/tasks/impl/nccl_task.h" @@ -19,7 +20,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { TEST_CASE("NCCL task spawns successfully") { std::vector fake_args = - make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/0_n); + make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/1_n); int fake_argc = fake_args.size(); char **fake_argv = fake_args.data(); @@ -28,6 +29,18 @@ TEST_CASE("NCCL task spawns successfully") { ControllerTaskResult result = manager.start_controller([](RealmContext &ctx) { + Realm::Machine::ProcessorQuery processor_query( + Realm::Machine::get_machine()); + processor_query.only_kind(Realm::Processor::TOC_PROC); + + Realm::Processor gpu_proc = processor_query.first(); + + DistributedFfHandle distributed_handle = + create_distributed_ff_handle( + ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true, + Realm::Event::NO_EVENT); DynamicNodeInvocation invocation{ /*inputs=*/{}, /*node_attrs=*/ @@ -38,12 +51,8 @@ TEST_CASE("NCCL task spawns successfully") { /*op_attrs=*/ TrainingOperationAttrs{ PCGOperatorAttrs{ - BroadcastAttrs{ - TensorDims{ - FFOrdered{ - 8_p, - }, - }, + ReplicateAttrs{ + /*replicate_degree=*/8_p, }, }, }, @@ -63,9 +72,10 @@ TEST_CASE("NCCL task spawns successfully") { Realm::Event event = spawn_nccl_task( ctx, - ctx.get_current_processor(), + gpu_proc, invocation, tensor_backing, + distributed_handle.at(gpu_proc), Realm::Event::NO_EVENT); event.wait();