Shamrock 2025.10.0
Astrophysical Code
Loading...
Searching...
No Matches
distributedDataComm.cpp
Go to the documentation of this file.
1// -------------------------------------------------------//
2//
3// SHAMROCK code for hydrodynamics
4// Copyright (c) 2021-2026 Timothée David--Cléris <tim.shamrock@proton.me>
5// SPDX-License-Identifier: CeCILL Free Software License Agreement v2.1
6// Shamrock is licensed under the CeCILL 2.1 License, see LICENSE for more information
7//
8// -------------------------------------------------------//
9
16
19#include "shambase/memory.hpp"
25#include "shamcmdopt/env.hpp"
26#include <memory>
27#include <vector>
28
29auto SPARSE_COMM_MODE = shamcmdopt::getenv_str_default_register(
30 "SPARSE_COMM_MODE", "new", "Sparse communication mode (new=with cache, old=without cache)");
31
32namespace {
33 struct SparseCommMode {
34 enum Mode { NEW, OLD };
35 };
36
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;
42 } else {
43 throw std::invalid_argument(
44 "Invalid sparse communication mode, valid modes are: new, old");
45 }
46 };
47
48 bool use_old_sparse_comm_mode = parse_sparse_comm_mode() == SparseCommMode::OLD;
49
50 bool warning_printed = false;
51} // namespace
52
53namespace shamalgs::collective {
54
55 namespace details {
56 struct DataTmp {
57 u64 sender;
58 u64 receiver;
59 u64 length;
61
62 SerializeSize get_ser_sz() {
63 return SerializeHelper::serialize_byte_size<u64>() * 3
64 + SerializeHelper::serialize_byte_size<u8>(length);
65 }
66 };
67
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)
71 -> std::map<std::pair<i32, i32>, SerializeHelper> {
72
73 StackEntry stack_loc{};
74
75 std::map<std::pair<i32, i32>, SerializeHelper> serializers;
76
77 for (auto &[key, vect] : send_data) {
78 SerializeSize byte_sz = SerializeHelper::serialize_byte_size<u64>(); // vec length
79 for (DataTmp &d : vect) {
80 byte_sz += d.get_ser_sz();
81 }
82 serializers.emplace(key, dev_sched);
83 serializers.at(key).allocate(byte_sz);
84 }
85
86 for (auto &[key, vect] : send_data) {
87 SerializeHelper &ser = serializers.at(key);
88 ser.write<u64>(vect.size());
89 for (DataTmp &d : vect) {
90 ser.write(d.sender);
91 ser.write(d.receiver);
92 ser.write(d.length);
93 ser.write_buf(d.data, d.length);
94 }
95 }
96
97 return serializers;
98 }
99
101 public:
102 i32 sender_rank, receiver_rank;
103 SerializeSize sz;
104 std::vector<std::reference_wrapper<DataTmp>> sources;
105 std::unique_ptr<SerializeHelper> serializer = {};
106 std::unique_ptr<sham::DeviceBuffer<u8>> send_buf = {};
107
108 void allocate_serializer(std::shared_ptr<sham::DeviceScheduler> dev_sched) {
109 serializer = std::make_unique<SerializeHelper>(dev_sched);
110 serializer->allocate(sz);
111 }
112
113 void write_sources() {
114 SerializeHelper &ser = shambase::get_check_ref(serializer);
115 ser.write<u64>(sources.size());
116 for (DataTmp &d : sources) {
117 ser.write(d.sender);
118 ser.write(d.receiver);
119 ser.write(d.length);
120 ser.write_buf(d.data, d.length);
121 }
122 }
123
124 void finalize_serializer() {
125 SerializeHelper &ser = shambase::get_check_ref(serializer);
126 send_buf = std::make_unique<sham::DeviceBuffer<u8>>(ser.finalize());
127 }
128 };
129
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> {
134
135 StackEntry stack_loc{};
136
137 std::vector<PrepareCommUtil> ret;
138
139 auto add_to_ret = [&](std::pair<i32, i32> key,
140 SerializeSize &byte_sz,
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()));
145 }
146
147 auto [sender_rank, receiver_rank] = key;
148
149 if (sources.size() > 0) {
150 PrepareCommUtil next{
151 .sender_rank = sender_rank,
152 .receiver_rank = receiver_rank,
153 .sz = byte_sz,
154 .sources = sources};
155 ret.push_back(std::move(next));
156 }
157
158 byte_sz = SerializeHelper::serialize_byte_size<u64>(); // vec length
159 sources = {};
160 };
161
162 for (auto &[key, vect] : send_data) {
163 SerializeSize byte_sz = SerializeHelper::serialize_byte_size<u64>(); // vec length
164 std::vector<std::reference_wrapper<DataTmp>> sources = {};
165
166 for (DataTmp &d : vect) {
167 std::reference_wrapper<DataTmp> d_ref = d;
168 auto dbyte_sz = d.get_ser_sz();
169
170 if ((dbyte_sz + byte_sz).get_total_size() > max_comm_size) {
171 add_to_ret(key, byte_sz, sources);
172 // logger::raw_ln("comm split at", d.sender, d.receiver, d.length);
173 }
174
175 // logger::raw_ln(
176 // "add to sources", dbyte_sz.get_total_size(), byte_sz.get_total_size());
177
178 byte_sz += d.get_ser_sz();
179 sources.push_back(d_ref);
180 }
181
182 add_to_ret(key, byte_sz, sources);
183 }
184
185 for (auto &c : ret) {
186 // logger::raw_ln(
187 // "allocate serializer", c.sender_rank, c.receiver_rank,
188 // c.sz.get_total_size());
189 c.allocate_serializer(dev_sched);
190 }
191
192 for (auto &c : ret) {
193 c.write_sources();
194 }
195
196 for (auto &c : ret) {
197 c.finalize_serializer();
198 }
199
200 return ret;
201 }
202
203 } // namespace details
204
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) {
211
212 StackEntry stack_loc{};
213
214 using namespace shambase;
215 using DataTmp = details::DataTmp;
216
217 // prepare map
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)};
221
222 send_data[key].push_back(
223 DataTmp{
224 .sender = sender, .receiver = receiver, .length = buf.get_size(), .data = buf});
225 });
226
227 // serialize together similar communications
228 std::map<std::pair<i32, i32>, SerializeHelper> serializers
229 = details::serialize_group_data(dev_sched, send_data);
230
231 // recover bufs from serializers
232 std::map<std::pair<i32, i32>, std::unique_ptr<sham::DeviceBuffer<u8>>> send_bufs;
233 {
234 NamedStackEntry stack_loc2{"recover bufs"};
235 for (auto &[key, ser] : serializers) {
236 send_bufs[key] = std::make_unique<sham::DeviceBuffer<u8>>(ser.finalize());
237 }
238 }
239
240 // prepare payload
241 std::vector<SendPayload> send_payoad;
242 {
243 NamedStackEntry stack_loc2{"prepare payload"};
244 for (auto &[key, buf] : send_bufs) {
245 send_payoad.push_back(
246 {.receiver_rank = key.second,
247 .payload = std::make_unique<shamcomm::CommunicationBuffer>(
248 shambase::extract_pointer(buf), dev_sched)});
249 }
250 }
251
252 // sparse comm
253 std::vector<RecvPayload> recv_payload;
254
255 if (comm_table) {
256 sparse_comm_c(dev_sched, send_payoad, recv_payload, *comm_table);
257 } else {
258 base_sparse_comm(dev_sched, send_payoad, recv_payload);
259 }
260
261 // make serializers from recv buffs
262 struct RecvPayloadSer {
263 i32 sender_ranks;
264 SerializeHelper ser;
265 };
266
267 std::vector<RecvPayloadSer> recv_payload_bufs;
268
269 {
270 NamedStackEntry stack_loc2{"move payloads"};
271 for (RecvPayload &payload : recv_payload) {
272
273 shamcomm::CommunicationBuffer comm_buf = extract_pointer(payload.payload);
274
275 sham::DeviceBuffer<u8> buf
276 = shamcomm::CommunicationBuffer::convert_usm(std::move(comm_buf));
277
278 recv_payload_bufs.push_back(
279 RecvPayloadSer{
280 .sender_ranks = payload.sender_ranks,
281 .ser = SerializeHelper(dev_sched, std::move(buf))});
282 }
283 }
284
285 {
286 NamedStackEntry stack_loc2{"split recv comms"};
287 // deserialize into the shared distributed data
288 for (RecvPayloadSer &recv : recv_payload_bufs) {
289 u64 cnt_obj;
290 recv.ser.load(cnt_obj);
291 for (u32 i = 0; i < cnt_obj; i++) {
292 u64 sender, receiver, length;
293
294 recv.ser.load(sender);
295 recv.ser.load(receiver);
296 recv.ser.load(length);
297
298 { // check correctness ranks
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");
304 }
305 }
306
307 auto it = recv_distrib_data.add_obj(
308 sender, receiver, sham::DeviceBuffer<u8>(length, dev_sched));
309
310 recv.ser.load_buf(it->second, length);
311 }
312 }
313 }
314 }
315
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,
321 DDSCommCache &cache,
322 std::optional<SparseCommTable> comm_table,
323 size_t max_comm_size) {
324
325 if (use_old_sparse_comm_mode) {
326 if (shamcomm::world_rank() == 0 && !warning_printed) {
327 logger::warn_ln("SparseComm", "using old sparse communication mode");
328 warning_printed = true;
329 }
330 return distributed_data_sparse_comm_old(
331 dev_sched, send_distrib_data, recv_distrib_data, rank_getter, comm_table);
332 }
333
335
336 using namespace shambase;
337 using DataTmp = details::DataTmp;
338
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;
342 } else {
343 max_alloc_size = dev_sched->ctx->device->prop.max_mem_alloc_size_host;
344 }
345 max_alloc_size -= 1; // keep a bit of space for safety
346
347 if (max_alloc_size > max_comm_size) {
348 max_alloc_size = max_comm_size;
349 }
350
351 // prepare map
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)};
355
356 send_data[key].push_back(
357 DataTmp{
358 .sender = sender, .receiver = receiver, .length = buf.get_size(), .data = buf});
359 });
360
361 std::vector<details::PrepareCommUtil> prepared_comms
362 = details::serialize_group_data_max_size(dev_sched, send_data, max_comm_size);
363
364 std::vector<shamalgs::collective::CommMessageInfo> messages_send;
365 std::vector<std::unique_ptr<sham::DeviceBuffer<u8>>> data_send;
366
367 for (auto &cms : prepared_comms) {
368
369 auto sender = cms.sender_rank;
370 auto receiver = cms.receiver_rank;
371 auto size = shambase::get_check_ref(cms.send_buf).get_size();
372
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,
381 });
382
383 data_send.push_back(std::move(cms.send_buf));
384 }
385
386 shamalgs::collective::CommTable comm_table2
387 = shamalgs::collective::build_sparse_exchange_table(messages_send, max_alloc_size);
388
389 if (dev_sched->ctx->device->mpi_prop.is_mpi_direct_capable) {
390 cache.set_sizes<sham::device>(
391 dev_sched, comm_table2.send_total_sizes, comm_table2.recv_total_sizes);
392 } else {
393 cache.set_sizes<sham::host>(
394 dev_sched, comm_table2.send_total_sizes, comm_table2.recv_total_sizes);
395 }
396
397 if (comm_table2.messages_send.size() != data_send.size()) {
398 std::vector<size_t> tmp1{};
399 for (size_t i = 0; i < data_send.size(); i++) {
400 tmp1.push_back(comm_table2.messages_send[i].message_size);
401 }
402
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());
406 }
407
408 throw make_except_with_loc<std::runtime_error>(
409 shambase::format("message send mismatch : {} != {}", tmp1, tmp2));
410 }
411
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++) {
415 tmp1.push_back(comm_table2.messages_send[i].message_size);
416 }
417
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);
421 }
422 throw make_except_with_loc<std::runtime_error>(
423 shambase::format("message send mismatch : {} != {}", tmp1, tmp2));
424 }
425
426 for (size_t i = 0; i < comm_table2.messages_send.size(); i++) {
427 auto &msg_info = comm_table2.messages_send[i];
428 auto offset_info = shambase::get_check_ref(msg_info.message_bytebuf_offset_send);
429 auto &buf_src = shambase::get_check_ref(data_send.at(i));
430
431 SHAM_ASSERT(buf_src.get_size() == msg_info.message_size);
432
433 cache.send_cache_write_buf_at(offset_info.buf_id, offset_info.data_offset, buf_src);
434 }
435
436 if (dev_sched->ctx->device->mpi_prop.is_mpi_direct_capable) {
437 shamalgs::collective::sparse_exchange<sham::device>(
438 dev_sched,
439 cache.get_cache1<sham::device>(),
440 cache.get_cache2<sham::device>(),
441 comm_table2);
442 } else {
443 shamalgs::collective::sparse_exchange<sham::host>(
444 dev_sched,
445 cache.get_cache1<sham::host>(),
446 cache.get_cache2<sham::host>(),
447 comm_table2);
448 }
449
450 // make serializers from recv buffs
451 struct RecvPayloadSer {
452 i32 sender_ranks;
453 SerializeHelper ser;
454 };
455
456 std::vector<RecvPayloadSer> recv_payload_bufs;
457
458 for (auto &msg : comm_table2.messages_recv) {
459
460 u64 size = msg.message_size;
461 i32 sender = msg.rank_sender;
462 i32 receiver = msg.rank_receiver;
463
464 auto offset_info = shambase::get_check_ref(msg.message_bytebuf_offset_recv);
465
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);
468
469 recv_payload_bufs.push_back(
470 RecvPayloadSer{
471 .sender_ranks = sender, .ser = SerializeHelper(dev_sched, std::move(recov))});
472 }
473
474 {
475 NamedStackEntry stack_loc2{"split recv comms"};
476 // deserialize into the shared distributed data
477 for (RecvPayloadSer &recv : recv_payload_bufs) {
478 u64 cnt_obj;
479 recv.ser.load(cnt_obj);
480 for (u32 i = 0; i < cnt_obj; i++) {
481 u64 sender, receiver, length;
482
483 recv.ser.load(sender);
484 recv.ser.load(receiver);
485 recv.ser.load(length);
486
487 { // check correctness ranks
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,
494 real_sender_rank));
495 }
496 }
497
498 auto it = recv_distrib_data.add_obj(
499 sender, receiver, sham::DeviceBuffer<u8>(length, dev_sched));
500
501 recv.ser.load_buf(it->second, length);
502 }
503 }
504 }
505 }
506
507} // namespace shamalgs::collective
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.
Definition assert.hpp:67
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.
@ host
Host memory.
@ device
Device memory.
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...
Definition memory.hpp:110
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
Definition memory.hpp:227
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,...
Definition env.hpp:97
i32 world_rank()
Gives the rank of the current process in the MPI communicator.
Definition worldInfo.cpp:40
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.