Skip to content
Closed
Show file tree
Hide file tree
Changes from 14 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
11 changes: 1 addition & 10 deletions cpp/benchmarks/bench_comm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -317,16 +317,7 @@ int main(int argc, char** argv) {
set_current_rmm_resource(args.rmm_mr);

rmm::device_async_resource_ref mr = cudf::get_current_device_resource_ref();
BufferResource br{
mr,
PinnedMemoryResource::Disabled,
{},
std::chrono::milliseconds{1},
std::make_shared<rmm::cuda_stream_pool>(
16, rmm::cuda_stream::flags::non_blocking
),
Comment thread
nirandaperera marked this conversation as resolved.
stats
};
BufferResource br{stats, mr};

// Print benchmark/hardware info.
{
Expand Down
4 changes: 2 additions & 2 deletions cpp/benchmarks/bench_partition.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ static void BM_PartitionAndPack(benchmark::State& state) {

// Create a pool memory resource with 50% of GPU memory
rmm::mr::pool_memory_resource pool_mr{rmm::mr::cuda_memory_resource{}, pool_size};
rapidsmpf::BufferResource br{pool_mr};
rapidsmpf::BufferResource br{rapidsmpf::Statistics::disabled(), pool_mr};

// Create input table
auto table = create_int_table(num_rows, stream);
Expand Down Expand Up @@ -111,7 +111,7 @@ static void BM_PartitionAndPackCurrentImpl(benchmark::State& state) {

// Create a pool memory resource with 50% of GPU memory
rmm::mr::pool_memory_resource pool_mr{rmm::mr::cuda_memory_resource{}, pool_size};
rapidsmpf::BufferResource br{pool_mr};
rapidsmpf::BufferResource br{rapidsmpf::Statistics::disabled(), pool_mr};

// Create input table
auto table = create_int_table(num_rows, stream);
Expand Down
8 changes: 2 additions & 6 deletions cpp/benchmarks/bench_shuffle.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -548,15 +548,11 @@ int main(int argc, char** argv) {
// We're only going to measure the last run, so disable initially.
stats->disable();
rapidsmpf::BufferResource br{
stats,
stat_enabled_mr,
args.pinned_mem_disable ? rapidsmpf::PinnedMemoryResource::Disabled
: rapidsmpf::PinnedMemoryResource::make_if_available(),
std::move(memory_limits),
std::chrono::milliseconds{1},
std::make_shared<rmm::cuda_stream_pool>(
16, rmm::cuda_stream::flags::non_blocking
),
stats
std::move(memory_limits)
};

std::shared_ptr<rapidsmpf::Communicator> comm;
Expand Down
22 changes: 7 additions & 15 deletions cpp/benchmarks/streaming/bench_streaming_shuffle.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -324,7 +324,9 @@ int main(int argc, char** argv) {

// Initialize configuration options from environment variables.
rapidsmpf::config::Options options{rapidsmpf::config::get_environment_variables()};
auto progress_thread = std::make_shared<rapidsmpf::ProgressThread>();

auto stats = rapidsmpf::Statistics::create();
auto progress_thread = std::make_shared<rapidsmpf::ProgressThread>(stats);

std::shared_ptr<rapidsmpf::Communicator> comm;
if (args.comm_type == "mpi") {
Expand Down Expand Up @@ -364,20 +366,11 @@ int main(int argc, char** argv) {
memory_limits[rapidsmpf::MemoryType::DEVICE] = args.device_mem_limit_mb << 20;
}

auto stats = rapidsmpf::Statistics::create();

auto pinned_mr = args.pinned_mem_disable
? rapidsmpf::PinnedMemoryResource::Disabled
: rapidsmpf::PinnedMemoryResource::make_if_available();
auto br = std::make_shared<rapidsmpf::BufferResource>(
stat_enabled_mr,
pinned_mr,
std::move(memory_limits),
std::nullopt,
std::make_shared<rmm::cuda_stream_pool>(
16, rmm::cuda_stream::flags::non_blocking
),
stats
stats, stat_enabled_mr, pinned_mr, std::move(memory_limits)
Comment thread
nirandaperera marked this conversation as resolved.
Outdated
);

auto& log = *comm->logger();
Expand Down Expand Up @@ -411,7 +404,7 @@ int main(int argc, char** argv) {
for (std::uint64_t i = 0; i < total_num_runs; ++i) {
// Clear statistics before the last run so only the final run is reported.
if (i == total_num_runs - 1) {
ctx->statistics()->clear();
stats->clear();
}

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.

Does the ctx not advertise a statistics object any more? That seems wrong.

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.

It does, but its the same as stats in this scope

double const elapsed = run(ctx, comm, args, stream).count();
std::stringstream ss;
Expand Down Expand Up @@ -458,15 +451,14 @@ int main(int argc, char** argv) {
log.print(ss.str());
}

auto statistics = ctx->statistics();
if (args.enable_memory_profiler) {
log.print(statistics->report({
log.print(stats->report({
.mr = stat_enabled_mr,
.pinned_mr = pinned_mr,
.header = "Statistics (of the last run):",
}));
} else {
log.print(statistics->report({.header = "Statistics (of the last run):"}));
log.print(stats->report({.header = "Statistics (of the last run):"}));
}

if (!use_bootstrap) {
Expand Down
4 changes: 2 additions & 2 deletions cpp/benchmarks/streaming/ndsh/utils.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -157,15 +157,15 @@ create_context(
);

auto br = std::make_shared<BufferResource>(
statistics,
std::move(mr),
arguments.no_pinned_host_memory ? PinnedMemoryResource::Disabled
: PinnedMemoryResource::make_if_available(),
std::move(memory_limits),
arguments.periodic_spill,
std::make_shared<rmm::cuda_stream_pool>(
arguments.num_streams, rmm::cuda_stream::flags::non_blocking
),
statistics
)
);
auto environment = config::get_environment_variables();
environment["NUM_STREAMING_THREADS"] =
Expand Down
2 changes: 1 addition & 1 deletion cpp/examples/example_shuffle.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ int main(int argc, char** argv) {
// We will use the same stream, memory, and buffer resource throughout the example.
rmm::cuda_stream_view stream = cudf::get_default_stream();
rmm::device_async_resource_ref mr = cudf::get_current_device_resource_ref();
rapidsmpf::BufferResource br{mr};
rapidsmpf::BufferResource br{stats, mr};

// As input data, we use a helper function from the benchmark suite. It creates a
// random cudf table with 2 columns and 100 rows. In this example, each MPI rank
Expand Down
4 changes: 3 additions & 1 deletion cpp/include/rapidsmpf/bootstrap/ucxx.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,9 @@ namespace bootstrap {
* @throws std::runtime_error if initialization fails.
*
* @code
* auto progress = std::make_shared<rapidsmpf::ProgressThread>();
* auto progress = std::make_shared<rapidsmpf::ProgressThread>(
* rapidsmpf::Statistics::disabled()
* );
* auto comm = rapidsmpf::bootstrap::create_ucxx_comm(progress);
* comm->logger().print("Hello from rank " + std::to_string(comm->rank()));
* @endcode
Expand Down
45 changes: 41 additions & 4 deletions cpp/include/rapidsmpf/communicator/communicator.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <rapidsmpf/error.hpp>
#include <rapidsmpf/memory/buffer.hpp>
#include <rapidsmpf/progress_thread.hpp>
#include <rapidsmpf/statistics.hpp>

/**
* @namespace rapidsmpf
Expand Down Expand Up @@ -403,7 +404,19 @@ class Communicator {
};

protected:
Communicator() = default;
/**
* @brief Construct the base Communicator with a progress thread.
*
* The communicator delegates `statistics()` to @p progress_thread so the
* two always agree on the current `Statistics` instance, including after
* `ProgressThread::set_statistics()` swaps it.
*
* @param progress_thread The progress thread for this communicator. Must
* not be null.
*
* @throws std::invalid_argument If @p progress_thread is null.
*/
explicit Communicator(std::shared_ptr<ProgressThread> progress_thread);

public:
virtual ~Communicator() noexcept = default;
Expand Down Expand Up @@ -614,18 +627,42 @@ class Communicator {

/**
* @brief Retrieves the progress thread associated with this communicator.
* @return Shared pointer to the progress thread.
* @return Shared pointer to the progress thread (never null).
*/
[[nodiscard]] virtual std::shared_ptr<ProgressThread> const&
progress_thread() const = 0;
[[nodiscard]] std::shared_ptr<ProgressThread> const&
progress_thread() const noexcept {
return progress_thread_;
}

/**
* @brief Retrieves the statistics instance associated with this communicator.
*
* Forwards to `progress_thread()->statistics()`, so the communicator and
* its progress thread always agree on the current instance — including
* after `ProgressThread::set_statistics()` swaps it. Satisfies the
* `StatisticsProvider` concept.
*
* @return Shared pointer to the statistics instance (never null).
*/
[[nodiscard]] std::shared_ptr<Statistics> statistics() const noexcept {
return progress_thread_->statistics();
}

/**
* @brief Provides a string representation of the communicator.
* @return A string describing the communicator.
*/
[[nodiscard]] virtual std::string str() const = 0;

private:
/// Progress thread owning this communicator's `Statistics` instance.
/// Never null after construction. Accessed only through `progress_thread()`
/// and `statistics()` so derived classes cannot bypass that contract.
std::shared_ptr<ProgressThread> progress_thread_;
};

static_assert(StatisticsProvider<Communicator>);

/// @brief Whether RapidsMPF was built with the UCXX Communicator.
#ifdef RAPIDSMPF_HAVE_UCXX
constexpr bool COMM_HAVE_UCXX = true;
Expand Down
10 changes: 1 addition & 9 deletions cpp/include/rapidsmpf/communicator/mpi.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -249,14 +249,6 @@ class MPI final : public Communicator {
return logger_;
}

/**
* @copydoc Communicator::progress_thread
*/
[[nodiscard]] std::shared_ptr<ProgressThread> const&
progress_thread() const override {
return progress_thread_;
}

/**
* @copydoc Communicator::str
*/
Expand All @@ -267,8 +259,8 @@ class MPI final : public Communicator {
Rank rank_;
Rank nranks_;
std::shared_ptr<Logger> logger_;
std::shared_ptr<ProgressThread> progress_thread_;
};

static_assert(StatisticsProvider<MPI>);

} // namespace rapidsmpf
10 changes: 1 addition & 9 deletions cpp/include/rapidsmpf/communicator/single.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -191,23 +191,15 @@ class Single final : public Communicator {
return logger_;
}

/**
* @copydoc Communicator::progress_thread
*/
[[nodiscard]] std::shared_ptr<ProgressThread> const&
progress_thread() const override {
return progress_thread_;
}

/**
* @copydoc Communicator::str
*/
[[nodiscard]] std::string str() const override;

private:
std::shared_ptr<Logger> logger_;
std::shared_ptr<ProgressThread> progress_thread_;
};

static_assert(StatisticsProvider<Single>);

} // namespace rapidsmpf
11 changes: 2 additions & 9 deletions cpp/include/rapidsmpf/communicator/ucxx.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -293,14 +293,6 @@ class UCXX final : public Communicator {
return logger_;
}

/**
* @copydoc Communicator::progress_thread
*/
[[nodiscard]] std::shared_ptr<ProgressThread> const&
progress_thread() const override {
return progress_thread_;
}

/**
* @copydoc Communicator::str
*/
Expand Down Expand Up @@ -339,12 +331,13 @@ class UCXX final : public Communicator {
std::shared_ptr<SharedResources> shared_resources_;
config::Options options_;
std::shared_ptr<Logger> logger_;
std::shared_ptr<ProgressThread> progress_thread_;

std::shared_ptr<::ucxx::Endpoint> get_endpoint(Rank rank);
void progress_worker();
};

} // namespace ucxx

static_assert(StatisticsProvider<ucxx::UCXX>);

} // namespace rapidsmpf
12 changes: 7 additions & 5 deletions cpp/include/rapidsmpf/memory/buffer_resource.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,8 @@ class BufferResource {
* If pinned-host memory is disabled, available pinned-host memory is always reported
* as zero regardless of the configured limit.
*
* @param statistics The statistics instance to use. Pass `Statistics::disabled()`
* to opt out of statistics collection.
* @param device_mr The RMM device memory resource used for device allocations.
* @param pinned_mr The pinned host memory resource used for `MemoryType::PINNED_HOST`
* allocations. If disabled, pinned host allocations are unavailable regardless of
Expand All @@ -80,16 +82,15 @@ class BufferResource {
* periodic spill checking is performed.
* @param stream_pool Pool of CUDA streams used throughout RapidsMPF for operations
* that do not take an explicit CUDA stream.
* @param statistics The statistics instance to use (disabled by default).
*/
BufferResource(
std::shared_ptr<Statistics> statistics,
cuda::mr::any_resource<cuda::mr::device_accessible> device_mr,
std::optional<PinnedMemoryResource> pinned_mr = PinnedMemoryResource::Disabled,
std::unordered_map<MemoryType, std::int64_t> memory_limits = {},
std::optional<Duration> periodic_spill_check = std::chrono::milliseconds{1},
std::shared_ptr<rmm::cuda_stream_pool> stream_pool = std::make_shared<
rmm::cuda_stream_pool>(16, rmm::cuda_stream::flags::non_blocking),
std::shared_ptr<Statistics> statistics = Statistics::disabled()
rmm::cuda_stream_pool>(16, rmm::cuda_stream::flags::non_blocking)
);

/**
Expand All @@ -101,15 +102,16 @@ class BufferResource {
*
* @param mr A device-accessible RMM memory resource.
* @param options Configuration options.
* @param statistics The statistics instance to use (disabled by default).
* @param statistics The statistics instance to use. Pass `Statistics::disabled()`
* to opt out of statistics collection.
*
* @return A shared pointer to a BufferResource instance configured according to the
* options.
*/
static std::shared_ptr<BufferResource> from_options(
cuda::mr::any_resource<cuda::mr::device_accessible> mr,
config::Options options,
std::shared_ptr<Statistics> statistics = Statistics::disabled()
std::shared_ptr<Statistics> statistics
);

~BufferResource() noexcept = default;
Expand Down
Loading