Shamrock 2025.10.0
Astrophysical Code
Loading...
Searching...
No Matches
PatchDataLayer.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
17
21#include "shambase/string.hpp"
27#include <type_traits>
28#include <variant>
29#include <vector>
30
31namespace shamrock::patch {
32
33 PatchDataLayer PatchDataLayer::mock_patchdata(
34 u64 seed, u32 obj_cnt, const std::shared_ptr<PatchDataLayerLayout> &pdl) {
35 PatchDataLayer pdat{pdl};
36
37 pdat.fields.clear();
38
39 pdat.pdl().for_each_field_any([&](auto &field) {
40 using f_t = typename std::remove_reference<decltype(field)>::type;
41 using base_t = typename f_t::field_T;
42
43 pdat.fields.push_back(
44 var_t{PatchDataField<base_t>::mock_field(seed, obj_cnt, field.name, field.nvar)});
45 });
46
47 return pdat;
48 }
49
50 PatchDataLayer::PatchDataLayer(const PatchDataLayer &other) : pdl_ptr(other.get_layout_ptr()) {
51
52 NamedStackEntry stack_loc{"PatchDataLayer::copy_constructor", true};
53
54 for (auto &field_var : other.fields) {
55
56 field_var.visit([&](auto &field) {
57 using base_t = typename std::remove_reference<decltype(field)>::type::Field_type;
58 fields.emplace_back(PatchDataField<base_t>(field));
59 });
60 };
61 }
62
63 void PatchDataLayer::init_fields() {
64
65 pdl().for_each_field_any([&](auto &field) {
66 using f_t = typename std::remove_reference<decltype(field)>::type;
67 using base_t = typename f_t::field_T;
68
69 fields.push_back(var_t{PatchDataField<base_t>(field.name, field.nvar)});
70 });
71 }
72
73 void PatchDataLayer::extract_element(u32 pidx, PatchDataLayer &out_pdat) {
74 StackEntry stack_loc{};
75
76 for (u32 idx = 0; idx < fields.size(); idx++) {
77
78 std::visit(
79 [&](auto &field, auto &out_field) {
80 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
81 using t2 =
82 typename std::remove_reference<decltype(out_field)>::type::Field_type;
83
84 if constexpr (std::is_same<t1, t2>::value) {
85 field.extract_element(pidx, out_field);
86 } else {
88 }
89 },
90 fields[idx].value,
91 out_pdat.fields[idx].value);
92 }
93 }
94
95 void PatchDataLayer::extract_elements(
96 const sham::DeviceBuffer<u32> &idxs, PatchDataLayer &out_pdat) {
97 StackEntry stack_loc{};
98
99 for (u32 idx = 0; idx < fields.size(); idx++) {
100
101 std::visit(
102 [&](auto &field, auto &out_field) {
103 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
104 using t2 =
105 typename std::remove_reference<decltype(out_field)>::type::Field_type;
106
107 if constexpr (std::is_same<t1, t2>::value) {
108 field.extract_elements(idxs, out_field);
109 } else {
111 }
112 },
113 fields[idx].value,
114 out_pdat.fields[idx].value);
115 }
116 }
117
118 void PatchDataLayer::insert_elements(const PatchDataLayer &pdat) {
119
120 StackEntry stack_loc{};
121
122 for (u32 idx = 0; idx < fields.size(); idx++) {
123
124 std::visit(
125 [&](auto &field, auto &out_field) {
126 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
127 using t2 =
128 typename std::remove_reference<decltype(out_field)>::type::Field_type;
129
130 if constexpr (std::is_same<t1, t2>::value) {
131 field.insert(out_field);
132 } else {
134 }
135 },
136 fields[idx].value,
137 pdat.fields[idx].value);
138 }
139 }
140
141 void PatchDataLayer::overwrite(PatchDataLayer &pdat, u32 obj_cnt) {
142 StackEntry stack_loc{};
143
144 for (u32 idx = 0; idx < fields.size(); idx++) {
145
146 std::visit(
147 [&](auto &field, auto &out_field) {
148 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
149 using t2 =
150 typename std::remove_reference<decltype(out_field)>::type::Field_type;
151
152 if constexpr (std::is_same<t1, t2>::value) {
153 field.overwrite(out_field, obj_cnt);
154 } else {
156 }
157 },
158 fields[idx].value,
159 pdat.fields[idx].value);
160 }
161 }
162
163 void PatchDataLayer::resize(u32 new_obj_cnt) {
164
165 for (auto &field_var : fields) {
166 field_var.visit([&](auto &field) {
167 field.resize(new_obj_cnt);
168 });
169 }
170 }
171
172 void PatchDataLayer::reserve(u32 new_obj_cnt) {
173
174 for (auto &field_var : fields) {
175 field_var.visit([&](auto &field) {
176 field.reserve(new_obj_cnt);
177 });
178 }
179 }
180
181 void PatchDataLayer::expand(u32 new_obj_cnt) {
182
183 for (auto &field_var : fields) {
184 field_var.visit([&](auto &field) {
185 field.expand(new_obj_cnt);
186 });
187 }
188 }
189
190 void PatchDataLayer::index_remap(sycl::buffer<u32> &index_map, u32 len) {
191
192 sham::DeviceBuffer<u32> dev_index_map(
193 index_map, len, shamsys::instance::get_compute_scheduler_ptr());
194
195 for (auto &field_var : fields) {
196 field_var.visit([&](auto &field) {
197 field.index_remap(dev_index_map, len);
198 });
199 }
200 }
201
202 void PatchDataLayer::index_remap_resize(sycl::buffer<u32> &index_map, u32 len) {
203 sham::DeviceBuffer<u32> dev_index_map(
204 index_map, len, shamsys::instance::get_compute_scheduler_ptr());
205
206 for (auto &field_var : fields) {
207 field_var.visit([&](auto &field) {
208 field.index_remap_resize(dev_index_map, len);
209 });
210 }
211 }
212
214 for (auto &field_var : fields) {
215 field_var.visit([&](auto &field) {
216 field.index_remap_resize(index_map, len);
217 });
218 }
219 }
220
221 void PatchDataLayer::keep_ids(sycl::buffer<u32> &index_map, u32 len) {
222 index_remap_resize(index_map, len);
223 }
224
225 void PatchDataLayer::keep_ids(sham::DeviceBuffer<u32> &index_map, u32 len) {
226 index_remap_resize(index_map, len);
227 }
228
230 for (auto &field_var : fields) {
231 field_var.visit([&](auto &field) {
232 field.remove_ids(indexes, len);
233 });
234 }
235 }
236
237 void PatchDataLayer::append_subset_to(
238 sycl::buffer<u32> &idxs_buf, u32 sz, PatchDataLayer &pdat) {
239 StackEntry stack_loc{};
240
241 for (u32 idx = 0; idx < fields.size(); idx++) {
242
243 std::visit(
244 [&](auto &field, auto &out_field) {
245 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
246 using t2 =
247 typename std::remove_reference<decltype(out_field)>::type::Field_type;
248
249 if constexpr (std::is_same<t1, t2>::value) {
250 field.append_subset_to(idxs_buf, sz, out_field);
251 } else {
253 }
254 },
255 fields[idx].value,
256 pdat.fields[idx].value);
257 }
258 }
259
260 void PatchDataLayer::append_subset_to(const std::vector<u32> &idxs, PatchDataLayer &pdat) {
261 StackEntry stack_loc{};
262
263 for (u32 idx = 0; idx < fields.size(); idx++) {
264
265 std::visit(
266 [&](auto &field, auto &out_field) {
267 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
268 using t2 =
269 typename std::remove_reference<decltype(out_field)>::type::Field_type;
270
271 if constexpr (std::is_same<t1, t2>::value) {
272 field.append_subset_to(idxs, out_field);
273 } else {
275 }
276 },
277 fields[idx].value,
278 pdat.fields[idx].value);
279 }
280 }
281
282 void PatchDataLayer::append_subset_to(
283 const sham::DeviceBuffer<u32> &idxs_buf, u32 sz, PatchDataLayer &pdat) const {
284 StackEntry stack_loc{};
285
286 for (u32 idx = 0; idx < fields.size(); idx++) {
287
288 std::visit(
289 [&](auto &field, auto &out_field) {
290 using t1 = typename std::remove_reference<decltype(field)>::type::Field_type;
291 using t2 =
292 typename std::remove_reference<decltype(out_field)>::type::Field_type;
293
294 if constexpr (std::is_same<t1, t2>::value) {
295 field.append_subset_to(idxs_buf, sz, out_field);
296 } else {
298 "Mismatch in layout\n source layout = {}\n dest layout = {}",
299 pdl().get_description_str(),
300 pdat.pdl().get_description_str()));
301 }
302 },
303 fields[idx].value,
304 pdat.fields[idx].value);
305 }
306 }
307
308 void PatchDataLayer::serialize_buf(shamalgs::SerializeHelper &serializer) {
309 StackEntry stack_loc{};
310 for_each_field_any([&](auto &f) {
311 f.serialize_buf(serializer);
312 });
313 }
314
315 shamalgs::SerializeSize PatchDataLayer::serialize_buf_byte_size() {
316 shamalgs::SerializeSize sum{};
317 for_each_field_any([&](auto &f) {
318 sum += f.serialize_buf_byte_size();
319 });
320 return sum;
321 }
322
323 PatchDataLayer PatchDataLayer::deserialize_buf(
324 shamalgs::SerializeHelper &serializer, const std::shared_ptr<PatchDataLayerLayout> &pdl) {
325 StackEntry stack_loc{};
326
327 return PatchDataLayer{pdl, [&](auto &pdat_fields) {
328 pdl->for_each_field_any([&](auto &field) {
329 using f_t =
330 typename std::remove_reference<decltype(field)>::type;
331 using base_t = typename f_t::field_T;
332
333 pdat_fields.push_back(
335 serializer, field.name, field.nvar)});
336 });
337 }};
338 }
339
340 void PatchDataLayer::fields_raz() {
341 for_each_field_any([&](auto &f) {
342 f.field_raz();
343 });
344 }
345
347
348 bool is_empty = fields.empty();
349
350 if (!is_empty) {
351 return fields[0].visit_return([](const auto &field) {
352 return field.get_obj_cnt();
353 });
354 }
355
357 "this PatchDataLayer does not contain any fields");
358 }
359
361 u64 sum = 0;
362
363 for (auto &field_var : fields) {
364
365 field_var.visit([&](auto &field) {
366 sum += field.memsize();
367 });
368 }
369
370 return sum;
371 }
372
374 for (auto &field_var : fields) {
375 field_var.visit([&](auto &field) {
376 field.synchronize_buf();
377 });
378 }
379 }
380
382 u32 cnt = get_obj_cnt();
383 for (auto &field_var : fields) {
384 field_var.visit([&](auto &field) {
385 if (field.get_obj_cnt() != cnt) {
386 throw shambase::make_except_with_loc<std::runtime_error>("mismatch in obj cnt");
387 }
388 });
389 }
390 }
391
393 StackEntry stack_loc{};
394
395 bool ret = false;
396
397 for (auto &field_var : fields) {
398 field_var.visit([&](auto &field) {
399 if (field.has_nan()) {
400 ret = true;
401 }
402 });
403 }
404 return ret;
405 }
406
408 StackEntry stack_loc{};
409
410 bool ret = false;
411
412 for (auto &field_var : fields) {
413 field_var.visit([&](auto &field) {
414 if (field.has_inf()) {
415 ret = true;
416 }
417 });
418 }
419 return ret;
420 }
421
423 StackEntry stack_loc{};
424
425 bool ret = false;
426
427 for (auto &field_var : fields) {
428 field_var.visit([&](auto &field) {
429 if (field.has_nan_or_inf()) {
430 ret = true;
431 }
432 });
433 }
434 return ret;
435 }
436
437 template<class T>
438 void PatchDataLayer::split_patchdata(
439 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
440 std::array<T, 8> min_box,
441 std::array<T, 8> max_box) {
442
443 StackEntry stack_loc{};
444
445 PatchDataField<T> &main_field = fields[0].get_if_ref_throw<T>();
446
447 // auto get_vec_idx = [&](T vmin, T vmax) -> std::vector<u32> {
448 // return main_field.get_elements_with_range(
449 // [&](T val, T vmin, T vmax) {
450 // return Patch::is_in_patch_converted(val, vmin, vmax);
451 // },
452 // vmin,
453 // vmax);
454 // };
455
456 auto get_vec_idx = [&](T vmin, T vmax) -> std::vector<u32> {
457 return main_field.get_ids_vec_where(
458 [&](const auto &acc, u32 idx, T vmin, T vmax) {
459 T val = acc[idx];
460 return Patch::is_in_patch_converted(val, vmin, vmax);
461 },
462 vmin,
463 vmax);
464 };
465
466 std::vector<u32> idx_p0 = get_vec_idx(min_box[0], max_box[0]);
467 std::vector<u32> idx_p1 = get_vec_idx(min_box[1], max_box[1]);
468 std::vector<u32> idx_p2 = get_vec_idx(min_box[2], max_box[2]);
469 std::vector<u32> idx_p3 = get_vec_idx(min_box[3], max_box[3]);
470 std::vector<u32> idx_p4 = get_vec_idx(min_box[4], max_box[4]);
471 std::vector<u32> idx_p5 = get_vec_idx(min_box[5], max_box[5]);
472 std::vector<u32> idx_p6 = get_vec_idx(min_box[6], max_box[6]);
473 std::vector<u32> idx_p7 = get_vec_idx(min_box[7], max_box[7]);
474
475 u32 el_cnt_new = idx_p0.size() + idx_p1.size() + idx_p2.size() + idx_p3.size()
476 + idx_p4.size() + idx_p5.size() + idx_p6.size() + idx_p7.size();
477
478 if (get_obj_cnt() != el_cnt_new) {
479
480 logger::err_ln(
481 "PatchData",
482 "error in patchdata split, the new element count doesn't match the old one");
483
484 logger::err_ln("PatchData", min_box[0], max_box[0]);
485 logger::err_ln("PatchData", min_box[1], max_box[1]);
486 logger::err_ln("PatchData", min_box[2], max_box[2]);
487 logger::err_ln("PatchData", min_box[3], max_box[3]);
488 logger::err_ln("PatchData", min_box[4], max_box[4]);
489 logger::err_ln("PatchData", min_box[5], max_box[5]);
490 logger::err_ln("PatchData", min_box[6], max_box[6]);
491 logger::err_ln("PatchData", min_box[7], max_box[7]);
492
493 T vmin = sham::min(min_box[0], min_box[1]);
494 vmin = sham::min(vmin, min_box[2]);
495 vmin = sham::min(vmin, min_box[3]);
496 vmin = sham::min(vmin, min_box[4]);
497 vmin = sham::min(vmin, min_box[5]);
498 vmin = sham::min(vmin, min_box[6]);
499 vmin = sham::min(vmin, min_box[7]);
500
501 T vmax = sham::max(max_box[0], max_box[1]);
502 vmax = sham::max(vmax, max_box[2]);
503 vmax = sham::max(vmax, max_box[3]);
504 vmax = sham::max(vmax, max_box[4]);
505 vmax = sham::max(vmax, max_box[5]);
506 vmax = sham::max(vmax, max_box[6]);
507 vmax = sham::max(vmax, max_box[7]);
508
509 main_field.check_err_range(
510 [&](T val, T vmin, T vmax) {
511 return Patch::is_in_patch_converted(val, vmin, vmax);
512 },
513 vmin,
514 vmax);
515 }
516
517 // TODO create a extract subpatch function
518
519 append_subset_to(idx_p0, pdats[0].get());
520 append_subset_to(idx_p1, pdats[1].get());
521 append_subset_to(idx_p2, pdats[2].get());
522 append_subset_to(idx_p3, pdats[3].get());
523 append_subset_to(idx_p4, pdats[4].get());
524 append_subset_to(idx_p5, pdats[5].get());
525 append_subset_to(idx_p6, pdats[6].get());
526 append_subset_to(idx_p7, pdats[7].get());
527 }
528
529#ifndef DOXYGEN
530 template void PatchDataLayer::split_patchdata(
531 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
532 std::array<f32_3, 8> min_box,
533 std::array<f32_3, 8> max_box);
534 template void PatchDataLayer::split_patchdata(
535 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
536 std::array<f64_3, 8> min_box,
537 std::array<f64_3, 8> max_box);
538 template void PatchDataLayer::split_patchdata(
539 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
540 std::array<u32_3, 8> min_box,
541 std::array<u32_3, 8> max_box);
542 template void PatchDataLayer::split_patchdata(
543 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
544 std::array<u64_3, 8> min_box,
545 std::array<u64_3, 8> max_box);
546 template void PatchDataLayer::split_patchdata(
547 std::array<std::reference_wrapper<PatchDataLayer>, 8> pdats,
548 std::array<i64_3, 8> min_box,
549 std::array<i64_3, 8> max_box);
550#endif
551
552 bool operator==(PatchDataLayer &p1, PatchDataLayer &p2) {
553 bool check = true;
554
555 if (p1.fields.size() != p2.fields.size()) {
556 return false;
557 }
558
559 for (u32 idx = 0; idx < p1.fields.size(); idx++) {
560 bool ret = p1.fields[idx].visit_return([&](auto &pf1) -> bool {
561 using t1 = typename std::remove_reference<decltype(pf1)>::type::Field_type;
562
563 if (PatchDataField<t1> *pf2
564 = std::get_if<PatchDataField<t1>>(&p2.fields[idx].value)) {
565 return pf1.check_field_match(*pf2);
566 } else {
567 return false;
568 }
569 });
570
571 check = check && ret;
572 }
573
574 return check;
575 }
576
577} // namespace shamrock::patch
Header file for the patch struct and related function.
std::uint32_t u32
32 bit unsigned integer
std::uint64_t u64
64 bit unsigned integer
static PatchDataField deserialize_buf(shamalgs::SerializeHelper &serializer, std::string field_name, u32 nvar)
deserialize a field inverse of serialize_buf
std::vector< u32 > get_ids_vec_where(Lambdacd &&cd_true, Args &&...args)
Same function as.
A buffer allocated in USM (Unified Shared Memory).
PatchDataLayer container class, the layout is described in patchdata_layout.
void index_remap(sycl::buffer< u32 > &index_map, u32 len)
this function remaps the patchdatafield like so val[id] = val[index_map[id]] This function can be use...
bool has_inf()
check whether any field contains an infinite value
friend bool operator==(PatchDataLayer &p1, PatchDataLayer &p2)
Compare two layers field by field (defined in PatchDataLayer.cpp).
void check_field_obj_cnt_match()
check that all contained field have the same obj cnt
u32 get_obj_cnt() const
get the number of objects (particles) stored in this layer
void synchronize_buf()
synchronize the host/device buffers of all fields
void index_remap_resize(sycl::buffer< u32 > &index_map, u32 len)
this function remaps the patchdatafield like so val[id] = val[index_map[id]] This function can be use...
void remove_ids(const sham::DeviceBuffer< u32 > &indexes, u32 len)
remove some particles ids
bool has_nan()
check whether any field contains a NaN value
void extract_element(u32 pidx, PatchDataLayer &out_pdat)
extract particle at index pidx and insert it in the provided vectors
bool has_nan_or_inf()
check whether any field contains a NaN or infinite value
u64 memsize()
get the memory size in bytes used by all fields
This header file contains utility functions related to exception handling in the code.
void append_subset_to(const sham::DeviceBuffer< T > &buf, const sham::DeviceBuffer< u32 > &idxs_buf, u32 nvar, sham::DeviceBuffer< T > &buf_other, u32 start_enque)
Appends a subset of elements from one buffer to another.
ExcptTypes make_except_with_loc(std::string message, SourceLocation loc=SourceLocation{})
Create an exception with a message and a location.
This file contains the definition for the stacktrace related functionality.
shambase::details::NamedBasicStackEntry NamedStackEntry
Alias for shambase::details::NamedBasicStackEntry.
shambase::details::BasicStackEntry StackEntry
Alias for shambase::details::BasicStackEntry.
header file to manage sycl