98 std::shared_ptr<shamrock::patch::PatchDataLayerLayout> pdl_ptr;
116 inline std::shared_ptr<shamrock::patch::PatchDataLayerLayout> get_layout_ptr_old()
const {
128 void init_mpi_required_types();
130 void free_mpi_required_types();
133 const std::shared_ptr<shamrock::patch::PatchDataLayerLayout> &pdl_ptr,
139 std::string dump_status();
141 inline void update_local_load_value(std::function<
u64(shamrock::patch::Patch)> load_function) {
144 p.load_value = load_function(p);
149 template<
class vectype>
150 std::tuple<vectype, vectype> get_box_tranform();
152 template<
class vectype>
153 std::tuple<vectype, vectype> get_box_volume();
155 bool should_resize_box(
bool node_in);
164 template<
class vectype>
167 if (!pdl_old().check_main_field_type<vectype>()) {
168 std::invalid_argument(
169 std::string(
"the main field is not of the correct type to call this function\n")
170 +
"fct called : " + __PRETTY_FUNCTION__
171 +
"current patch data layout : " + pdl_old().get_description_str());
174 patch_data.sim_box.set_bounding_box<vectype>({bmin, bmax});
176 shamlog_debug_ln(
"PatchScheduler",
"box resized to :", bmin, bmax);
188 void make_patch_base_grid(std::array<u32, dim> patch_count);
196 template<
class vectype>
204 void check_patchdata_locality_correctness();
207 void dump_local_patches(std::string filename);
209 std::vector<std::unique_ptr<shamrock::patch::PatchDataLayer>> gather_data(
u32 rank);
234 void sync_build_LB(
bool global_patch_sync,
bool balance_load);
238 return get_sim_box().template get_patch_transform<vec>();
260 template<
class Function>
268 fct(patch_id, cur_p, pdat);
273 template<
class Function>
274 inline void for_each_patch(Function &&fct) {
282 fct(patch_id, cur_p);
287 inline void for_each_global_patch(
288 const std::function<
void(
const shamrock::patch::Patch &)> &fct) {
289 for (
const shamrock::patch::Patch &p :
patch_list.global) {
290 if (!p.is_err_mode()) {
296 inline void for_each_local_patch(
297 const std::function<
void(
const shamrock::patch::Patch &)> &fct) {
298 for (
const shamrock::patch::Patch &p :
patch_list.local) {
299 if (!p.is_err_mode()) {
305 inline void for_each_local_patchdata(
306 const std::function<
void(
const shamrock::patch::Patch &, shamrock::patch::PatchDataLayer &)>
308 for (
const shamrock::patch::Patch &p :
patch_list.local) {
309 if (!p.is_err_mode()) {
315 inline void for_each_local_patch_nonempty(
316 std::function<
void(
const shamrock::patch::Patch &)> fct) {
317 patch_data.for_each_patchdata([&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
318 shamrock::patch::Patch &cur_p
321 if ((!cur_p.
is_err_mode()) && (!pdat.is_empty())) {
327 inline u32 get_patch_rank_owner(
u64 patch_id) {
328 shamrock::patch::Patch &cur_p
333 inline void for_each_patchdata_nonempty(
334 std::function<
void(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &)> fct) {
335 patch_data.for_each_patchdata([&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
336 shamrock::patch::Patch &cur_p
339 if ((!cur_p.
is_err_mode()) && (!pdat.is_empty())) {
346 inline shambase::DistributedData<T> map_owned_patchdata(
347 std::function<T(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &pdat)> fct) {
348 shambase::DistributedData<T> ret;
350 using namespace shamrock::patch;
352 ret.
add_obj(id_patch, fct(cur_p, pdat));
359 inline shambase::DistributedData<T> distrib_data_local_to_all_simple(
360 shambase::DistributedData<T> &src) {
361 using namespace shamrock::patch;
365 return shamalgs::collective::fetch_all_simple<T, Patch>(
372 inline shambase::DistributedData<T> distrib_data_local_to_all_load_store(
373 shambase::DistributedData<T> &src) {
374 using namespace shamrock::patch;
376 return shamalgs::collective::fetch_all_storeload<T, Patch>(
383 inline shambase::DistributedData<T> map_owned_patchdata_fetch_simple(
384 std::function<T(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &pdat)> fct) {
385 shambase::DistributedData<T> ret;
387 using namespace shamrock::patch;
389 ret.
add_obj(id_patch, fct(cur_p, pdat));
392 return distrib_data_local_to_all_simple(ret);
396 inline shambase::DistributedData<T> map_owned_patchdata_fetch_load_store(
397 std::function<T(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &pdat)> fct) {
398 shambase::DistributedData<T> ret;
400 using namespace shamrock::patch;
402 ret.
add_obj(id_patch, fct(cur_p, pdat));
405 return distrib_data_local_to_all_load_store(ret);
409 inline shamrock::patch::PatchField<T> map_owned_to_patch_field_simple(
410 std::function<T(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &pdat)> fct) {
411 return shamrock::patch::PatchField<T>(map_owned_patchdata_fetch_simple(fct));
415 inline shamrock::patch::PatchField<T> map_owned_to_patch_field_load_store(
416 std::function<T(
const shamrock::patch::Patch, shamrock::patch::PatchDataLayer &pdat)> fct) {
417 return shamrock::patch::PatchField<T>(map_owned_patchdata_fetch_load_store(fct));
420 inline u64 get_rank_count() {
422 using namespace shamrock::patch;
425 num_obj += pdat.get_obj_cnt();
431 inline u64 get_total_obj_count() {
433 u64 part_cnt = get_rank_count();
434 return shamalgs::collective::allreduce_sum(part_cnt);
438 inline std::unique_ptr<sycl::buffer<T>> rankgather_field(
u32 field_idx) {
440 std::unique_ptr<sycl::buffer<T>> ret;
442 auto fd = pdl_old().get_field<T>(field_idx);
445 u64 num_obj = get_rank_count();
448 ret = std::make_unique<sycl::buffer<T>>(num_obj * nvar);
450 using namespace shamrock::patch;
454 using namespace shamalgs::memory;
455 using namespace shambase;
457 if (pdat.get_obj_cnt() > 0) {
458 write_with_offset_into(
459 shamsys::instance::get_compute_scheduler().get_queue(),
461 pdat.get_field<T>(field_idx).get_buf(),
463 pdat.get_obj_cnt() * nvar);
465 ptr += pdat.get_obj_cnt() * nvar;
495 template<
class Function,
class Pfield>
496 inline void compute_patch_field(Pfield &field, MPI_Datatype &dtype, Function &&lambda) {
497 field.local_nodes_value.resize(
patch_list.local.size());
501 shamrock::patch::Patch &cur_p =
patch_list.local[idx];
504 field.local_nodes_value[idx] = lambda(
505 shamsys::instance::get_compute_queue(),
511 field.build_global(dtype);
514 inline auto get_node_set_edge_patchdata_layer_refs() {
515 shamrock::solvergraph::NodeSetEdge<shamrock::solvergraph::PatchDataLayerRefs> node_set_edge(
516 [&](shamrock::solvergraph::PatchDataLayerRefs &edge) {
518 using namespace shamrock::patch;
519 for_each_patchdata_nonempty([&](Patch cur_p, PatchDataLayer &pdat) {
524 return std::make_shared<decltype(node_set_edge)>(std::move(node_set_edge));
533 std::vector<u64>
add_root_patches(std::vector<shamrock::patch::PatchCoord<3>> coords);
535 shamrock::patch::SimulationBoxInfo &get_sim_box() {
return patch_data.sim_box; }
537 nlohmann::json serialize_patch_metadata();
540 void split_patches(std::unordered_set<u64> split_rq);
541 void merge_patches(std::unordered_set<u64> merge_rq);
543 void set_patch_pack_values(std::unordered_set<u64> merge_rq);