30 "SPARSE_COMM_MODE",
"new",
"Sparse communication mode (new=with cache, old=without cache)");
33 struct SparseCommMode {
34 enum Mode { NEW, OLD };
37 constexpr auto parse_sparse_comm_mode = []() {
38 if (SPARSE_COMM_MODE ==
"new") {
39 return SparseCommMode::NEW;
40 }
else if (SPARSE_COMM_MODE ==
"old") {
41 return SparseCommMode::OLD;
43 throw std::invalid_argument(
44 "Invalid sparse communication mode, valid modes are: new, old");
48 bool use_old_sparse_comm_mode = parse_sparse_comm_mode() == SparseCommMode::OLD;
50 bool warning_printed =
false;
53namespace shamalgs::collective {
63 return SerializeHelper::serialize_byte_size<u64>() * 3
64 + SerializeHelper::serialize_byte_size<u8>(length);
68 auto serialize_group_data(
69 std::shared_ptr<sham::DeviceScheduler> dev_sched,
70 std::map<std::pair<i32, i32>, std::vector<DataTmp>> &send_data)
77 for (
auto &[key, vect] : send_data) {
78 SerializeSize byte_sz = SerializeHelper::serialize_byte_size<u64>();
79 for (DataTmp &d : vect) {
80 byte_sz += d.get_ser_sz();
82 serializers.emplace(key, dev_sched);
83 serializers.at(key).allocate(byte_sz);
86 for (
auto &[key, vect] : send_data) {
87 SerializeHelper &ser = serializers.at(key);
88 ser.write<
u64>(vect.size());
89 for (DataTmp &d : vect) {
91 ser.write(d.receiver);
93 ser.write_buf(d.data, d.length);
102 i32 sender_rank, receiver_rank;
104 std::vector<std::reference_wrapper<DataTmp>> sources;
105 std::unique_ptr<SerializeHelper> serializer = {};
106 std::unique_ptr<sham::DeviceBuffer<u8>> send_buf = {};
108 void allocate_serializer(std::shared_ptr<sham::DeviceScheduler> dev_sched) {
109 serializer = std::make_unique<SerializeHelper>(dev_sched);
110 serializer->allocate(sz);
113 void write_sources() {
115 ser.write<
u64>(sources.size());
118 ser.write(d.receiver);
120 ser.write_buf(d.data, d.length);
124 void finalize_serializer() {
126 send_buf = std::make_unique<sham::DeviceBuffer<u8>>(ser.finalize());
130 auto serialize_group_data_max_size(
131 std::shared_ptr<sham::DeviceScheduler> dev_sched,
132 std::map<std::pair<i32, i32>, std::vector<DataTmp>> &send_data,
133 u64 max_comm_size) -> std::vector<PrepareCommUtil> {
137 std::vector<PrepareCommUtil> ret;
139 auto add_to_ret = [&](std::pair<i32, i32> key,
141 std::vector<std::reference_wrapper<DataTmp>> &sources) {
142 if (byte_sz.get_total_size() > max_comm_size) {
144 shambase::format(
"comm size too large: {}", byte_sz.get_total_size()));
147 auto [sender_rank, receiver_rank] = key;
149 if (sources.size() > 0) {
150 PrepareCommUtil next{
151 .sender_rank = sender_rank,
152 .receiver_rank = receiver_rank,
155 ret.push_back(std::move(next));
158 byte_sz = SerializeHelper::serialize_byte_size<u64>();
162 for (
auto &[key, vect] : send_data) {
163 SerializeSize byte_sz = SerializeHelper::serialize_byte_size<u64>();
164 std::vector<std::reference_wrapper<DataTmp>> sources = {};
167 std::reference_wrapper<DataTmp> d_ref = d;
168 auto dbyte_sz = d.get_ser_sz();
170 if ((dbyte_sz + byte_sz).get_total_size() > max_comm_size) {
171 add_to_ret(key, byte_sz, sources);
178 byte_sz += d.get_ser_sz();
179 sources.push_back(d_ref);
182 add_to_ret(key, byte_sz, sources);
185 for (
auto &c : ret) {
189 c.allocate_serializer(dev_sched);
192 for (
auto &c : ret) {
196 for (
auto &c : ret) {
197 c.finalize_serializer();
205 void distributed_data_sparse_comm_old(
206 sham::DeviceScheduler_ptr dev_sched,
207 SerializedDDataComm &send_distrib_data,
208 SerializedDDataComm &recv_distrib_data,
209 std::function<
i32(
u64)> rank_getter,
210 std::optional<SparseCommTable> comm_table) {
214 using namespace shambase;
218 std::map<std::pair<i32, i32>, std::vector<DataTmp>> send_data;
219 send_distrib_data.for_each([&](
u64 sender,
u64 receiver, sham::DeviceBuffer<u8> &buf) {
220 std::pair<i32, i32> key = {rank_getter(sender), rank_getter(receiver)};
222 send_data[key].push_back(
224 .sender = sender, .receiver = receiver, .length = buf.
get_size(), .data = buf});
228 std::map<std::pair<i32, i32>, SerializeHelper> serializers
229 = details::serialize_group_data(dev_sched, send_data);
232 std::map<std::pair<i32, i32>, std::unique_ptr<sham::DeviceBuffer<u8>>> send_bufs;
235 for (
auto &[key, ser] : serializers) {
236 send_bufs[key] = std::make_unique<sham::DeviceBuffer<u8>>(ser.finalize());
241 std::vector<SendPayload> send_payoad;
244 for (
auto &[key, buf] : send_bufs) {
245 send_payoad.push_back(
246 {.receiver_rank = key.second,
247 .payload = std::make_unique<shamcomm::CommunicationBuffer>(
253 std::vector<RecvPayload> recv_payload;
256 sparse_comm_c(dev_sched, send_payoad, recv_payload, *comm_table);
258 base_sparse_comm(dev_sched, send_payoad, recv_payload);
262 struct RecvPayloadSer {
267 std::vector<RecvPayloadSer> recv_payload_bufs;
273 shamcomm::CommunicationBuffer comm_buf =
extract_pointer(payload.payload);
275 sham::DeviceBuffer<u8> buf
278 recv_payload_bufs.push_back(
280 .sender_ranks = payload.sender_ranks,
281 .ser = SerializeHelper(dev_sched, std::move(buf))});
288 for (RecvPayloadSer &recv : recv_payload_bufs) {
290 recv.ser.load(cnt_obj);
291 for (
u32 i = 0; i < cnt_obj; i++) {
292 u64 sender, receiver, length;
294 recv.ser.load(sender);
295 recv.ser.load(receiver);
296 recv.ser.load(length);
299 i32 supposed_sender_rank = rank_getter(sender);
300 i32 real_sender_rank = recv.sender_ranks;
301 if (supposed_sender_rank != real_sender_rank) {
302 throw make_except_with_loc<std::runtime_error>(
303 "the rank do not matches");
307 auto it = recv_distrib_data.add_obj(
308 sender, receiver, sham::DeviceBuffer<u8>(length, dev_sched));
310 recv.ser.load_buf(it->second, length);
316 void distributed_data_sparse_comm(
317 sham::DeviceScheduler_ptr dev_sched,
318 SerializedDDataComm &send_distrib_data,
319 SerializedDDataComm &recv_distrib_data,
320 std::function<
i32(
u64)> rank_getter,
322 std::optional<SparseCommTable> comm_table,
323 size_t max_comm_size) {
325 if (use_old_sparse_comm_mode) {
327 logger::warn_ln(
"SparseComm",
"using old sparse communication mode");
328 warning_printed =
true;
330 return distributed_data_sparse_comm_old(
331 dev_sched, send_distrib_data, recv_distrib_data, rank_getter, comm_table);
336 using namespace shambase;
339 size_t max_alloc_size;
340 if (dev_sched->ctx->device->mpi_prop.is_mpi_direct_capable) {
341 max_alloc_size = dev_sched->ctx->device->prop.max_mem_alloc_size_dev;
343 max_alloc_size = dev_sched->ctx->device->prop.max_mem_alloc_size_host;
347 if (max_alloc_size > max_comm_size) {
348 max_alloc_size = max_comm_size;
352 std::map<std::pair<i32, i32>, std::vector<DataTmp>> send_data;
353 send_distrib_data.for_each([&](
u64 sender,
u64 receiver, sham::DeviceBuffer<u8> &buf) {
354 std::pair<i32, i32> key = {rank_getter(sender), rank_getter(receiver)};
356 send_data[key].push_back(
358 .sender = sender, .receiver = receiver, .length = buf.
get_size(), .data = buf});
361 std::vector<details::PrepareCommUtil> prepared_comms
362 = details::serialize_group_data_max_size(dev_sched, send_data, max_comm_size);
364 std::vector<shamalgs::collective::CommMessageInfo> messages_send;
365 std::vector<std::unique_ptr<sham::DeviceBuffer<u8>>> data_send;
367 for (
auto &cms : prepared_comms) {
369 auto sender = cms.sender_rank;
370 auto receiver = cms.receiver_rank;
373 messages_send.push_back(
374 shamalgs::collective::CommMessageInfo{
375 .message_size = size,
376 .rank_sender = sender,
377 .rank_receiver = receiver,
378 .message_tag = std::nullopt,
379 .message_bytebuf_offset_send = std::nullopt,
380 .message_bytebuf_offset_recv = std::nullopt,
383 data_send.push_back(std::move(cms.send_buf));
386 shamalgs::collective::CommTable comm_table2
387 = shamalgs::collective::build_sparse_exchange_table(messages_send, max_alloc_size);
389 if (dev_sched->ctx->device->mpi_prop.is_mpi_direct_capable) {
398 std::vector<size_t> tmp1{};
399 for (
size_t i = 0; i < data_send.size(); i++) {
403 std::vector<size_t> tmp2{};
404 for (
size_t i = 0; i < data_send.size(); i++) {
405 tmp2.push_back(data_send[i]->get_size());
408 throw make_except_with_loc<std::runtime_error>(
409 shambase::format(
"message send mismatch : {} != {}", tmp1, tmp2));
412 if (comm_table2.
messages_send.size() != messages_send.size()) {
413 std::vector<size_t> tmp1{};
414 for (
size_t i = 0; i < comm_table2.
messages_send.size(); i++) {
418 std::vector<size_t> tmp2{};
419 for (
size_t i = 0; i < messages_send.size(); i++) {
420 tmp2.push_back(messages_send[i].message_size);
422 throw make_except_with_loc<std::runtime_error>(
423 shambase::format(
"message send mismatch : {} != {}", tmp1, tmp2));
426 for (
size_t i = 0; i < comm_table2.
messages_send.size(); i++) {
431 SHAM_ASSERT(buf_src.get_size() == msg_info.message_size);
433 cache.send_cache_write_buf_at(offset_info.buf_id, offset_info.data_offset, buf_src);
436 if (dev_sched->ctx->device->mpi_prop.is_mpi_direct_capable) {
437 shamalgs::collective::sparse_exchange<sham::device>(
443 shamalgs::collective::sparse_exchange<sham::host>(
451 struct RecvPayloadSer {
456 std::vector<RecvPayloadSer> recv_payload_bufs;
460 u64 size = msg.message_size;
461 i32 sender = msg.rank_sender;
462 i32 receiver = msg.rank_receiver;
466 sham::DeviceBuffer<u8> recov(size, dev_sched);
467 cache.recv_cache_read_buf_at(offset_info.buf_id, offset_info.data_offset, size, recov);
469 recv_payload_bufs.push_back(
471 .sender_ranks = sender, .ser = SerializeHelper(dev_sched, std::move(recov))});
477 for (RecvPayloadSer &recv : recv_payload_bufs) {
479 recv.ser.load(cnt_obj);
480 for (
u32 i = 0; i < cnt_obj; i++) {
481 u64 sender, receiver, length;
483 recv.ser.load(sender);
484 recv.ser.load(receiver);
485 recv.ser.load(length);
488 i32 supposed_sender_rank = rank_getter(sender);
489 i32 real_sender_rank = recv.sender_ranks;
490 if (supposed_sender_rank != real_sender_rank) {
491 throw make_except_with_loc<std::runtime_error>(shambase::format(
492 "the rank do not matches {} != {}",
493 supposed_sender_rank,
498 auto it = recv_distrib_data.add_obj(
499 sender, receiver, sham::DeviceBuffer<u8>(length, dev_sched));
501 recv.ser.load_buf(it->second, length);
std::uint32_t u32
32 bit unsigned integer
std::uint64_t u64
64 bit unsigned integer
std::int32_t i32
32 bit integer
#define SHAM_ASSERT(x)
Shorthand for SHAM_ASSERT_NAMED without a message.
A buffer allocated in USM (Unified Shared Memory).
size_t get_size() const
Gets the number of elements in the buffer.
static sham::DeviceBuffer< u8 > convert_usm(CommunicationBuffer &&buf)
destroy the buffer and recover the held object
This header file contains utility functions related to exception handling in the code.
T & get_check_ref(const std::unique_ptr< T > &ptr, SourceLocation loc=SourceLocation())
Takes a std::unique_ptr and returns a reference to the object it holds. It throws a std::runtime_erro...
ExcptTypes make_except_with_loc(std::string message, SourceLocation loc=SourceLocation{})
Create an exception with a message and a location.
auto extract_pointer(std::unique_ptr< T > &o, SourceLocation loc=SourceLocation()) -> T
extract content out of unique_ptr
std::string getenv_str_default_register(const char *env_var, std::string default_val, std::string desc)
Get the content of the environment variable if it exist and register it documentation,...
i32 world_rank()
Gives the rank of the current process in the MPI communicator.
This file contains the definition for the stacktrace related functionality.
#define __shamrock_stack_entry()
Macro to create a stack entry.
shambase::details::NamedBasicStackEntry NamedStackEntry
Alias for shambase::details::NamedBasicStackEntry.
shambase::details::BasicStackEntry StackEntry
Alias for shambase::details::BasicStackEntry.
std::vector< CommMessageInfo > messages_recv
Messages to recv.
std::vector< size_t > recv_total_sizes
Total size of the recv buffer.
std::vector< size_t > send_total_sizes
Total size of the send buffer.
std::vector< CommMessageInfo > messages_send
Messages to send.