Shamrock 2025.10.0
Astrophysical Code
Loading...
Searching...
No Matches
AMRGridRefinementHandler.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
18#include "shambase/memory.hpp"
20#include "shamcomm/logs.hpp"
26#include <stdexcept>
27#include <variant>
28
29template<class Tvec, class TgridVec>
30template<class UserAcc, class... T>
35 T &&...args) {
36
37 using namespace shamrock::patch;
38
39 u64 tot_refine = 0;
40 u64 tot_derefine = 0;
41
42 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
43 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
44
45 scheduler().for_each_patchdata_nonempty([&](Patch cur_p, PatchDataLayer &pdat) {
46 u64 id_patch = cur_p.id_patch;
47
48 // create the refine and derefine flags buffers
49 u32 obj_cnt = pdat.get_obj_cnt();
50
51 sham::DeviceBuffer<u32> refine_flags(obj_cnt, dev_sched);
52 sham::DeviceBuffer<u32> derefine_flags(obj_cnt, dev_sched);
53
54 {
55 sham::EventList depends_list;
56
57 UserAcc uacc(depends_list, id_patch, cur_p, pdat, args...);
58
59 auto refine_acc = refine_flags.get_write_access(depends_list);
60 auto derefine_acc = derefine_flags.get_write_access(depends_list);
61
62 // fill in the flags
63 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
64 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
65 bool flag_refine = false;
66 bool flag_derefine = false;
67 uacc.refine_criterion(gid.get_linear_id(), uacc, flag_refine, flag_derefine);
68
69 // This is just a safe guard to avoid this nonsensicall case
70 if (flag_refine && flag_derefine) {
71 flag_derefine = false;
72 }
73
74 refine_acc[gid] = (flag_refine) ? 1 : 0;
75 derefine_acc[gid] = (flag_derefine) ? 1 : 0;
76 });
77 });
78
79 sham::EventList resulting_events;
80 resulting_events.add_event(e);
81
82 refine_flags.complete_event_state(resulting_events);
83 derefine_flags.complete_event_state(resulting_events);
84
85 uacc.finalize(resulting_events, id_patch, cur_p, pdat, args...);
86 }
87
88 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
89 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
90
91 sham::EventList depends_list;
92 auto acc_min = buf_cell_min.get_read_access(depends_list);
93 auto acc_max = buf_cell_max.get_read_access(depends_list);
94 auto acc_merge_flag = derefine_flags.get_write_access(depends_list);
95
96 // keep only derefine flags on only if the eight cells want to merge and if they can
97 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
98 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
99 u32 id = gid.get_linear_id();
100
101 std::array<BlockCoord, split_count> blocks;
102 bool do_merge = true;
103
104 // This avoid the case where we are in the last block of the buffer to avoid the
105 // out-of-bound read
106 if (id + split_count <= obj_cnt) {
107 bool all_want_to_merge = true;
108
109 for (u32 lid = 0; lid < split_count; lid++) {
110 blocks[lid] = BlockCoord{acc_min[gid + lid], acc_max[gid + lid]};
111 all_want_to_merge = all_want_to_merge && acc_merge_flag[gid + lid];
112 }
113
114 do_merge = all_want_to_merge && BlockCoord::are_mergeable(blocks);
115
116 } else {
117 do_merge = false;
118 }
119
120 acc_merge_flag[gid] = do_merge;
121 });
122 });
123
124 buf_cell_min.complete_event_state(e);
125 buf_cell_max.complete_event_state(e);
126 derefine_flags.complete_event_state(e);
127
129 // refinement
131
132 // perform stream compactions on the refinement flags
133 auto buf_refine = shamalgs::numeric::stream_compact(dev_sched, refine_flags, obj_cnt);
134
135 shamlog_debug_ln(
136 "AMRGrid", "patch ", id_patch, "refine block count = ", buf_refine.get_size());
137
138 tot_refine += buf_refine.get_size();
139
140 // add the results to the map
141 dd_refine_list.add_obj(id_patch, std::move(buf_refine));
142
144 // derefinement
146
147 // perform stream compactions on the derefinement flags
148 auto buf_derefine = shamalgs::numeric::stream_compact(dev_sched, derefine_flags, obj_cnt);
149
150 shamlog_debug_ln(
151 "AMRGrid", "patch ", id_patch, "merge block count = ", buf_derefine.get_size());
152
153 tot_derefine += buf_derefine.get_size();
154
155 // add the results to the map
156 dd_derefine_list.add_obj(id_patch, std::move(buf_derefine));
157 });
158
159 logger::info_ln("AMRGrid", "on this process", tot_refine, "blocks were refined");
161 "AMRGrid", "on this process", tot_derefine * split_count, "blocks were derefined");
162}
163template<class Tvec, class TgridVec>
164template<class UserAcc>
167
168 using namespace shamrock::patch;
169
170 u64 sum_block_count = 0;
171
172 bool new_cell_were_added = false;
173
174 scheduler().for_each_patch_data([&](u64 id_patch, Patch cur_p, PatchDataLayer &pdat) {
175 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
176
177 u32 old_obj_cnt = pdat.get_obj_cnt();
178
179 sham::DeviceBuffer<u32> &refine_flags = dd_refine_list.get(id_patch);
180
181 if (refine_flags.get_size() > 0) {
182
183 // alloc memory for the new blocks to be created
184 pdat.expand(refine_flags.get_size() * (split_count - 1));
185
186 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
187 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
188
189 sham::EventList depends_list;
190 auto block_bound_low = buf_cell_min.get_write_access(depends_list);
191 auto block_bound_high = buf_cell_max.get_write_access(depends_list);
192 UserAcc uacc(depends_list, pdat);
193 auto index_to_ref = refine_flags.get_read_access(depends_list);
194
195 // Refine the block (set the positions) and fill the corresponding fields
196 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
197 u32 start_index_push = old_obj_cnt;
198
199 constexpr u32 new_splits = split_count - 1;
200
201 cgh.parallel_for(sycl::range<1>(refine_flags.get_size()), [=](sycl::item<1> gid) {
202 u32 tid = gid.get_linear_id();
203
204 u32 idx_to_refine = index_to_ref[tid];
205
206 // gen splits coordinates
207 BlockCoord cur_block{
208 block_bound_low[idx_to_refine], block_bound_high[idx_to_refine]};
209
210 std::array<BlockCoord, split_count> block_coords
211 = BlockCoord::get_split(cur_block.bmin, cur_block.bmax);
212
213 // generate index for the refined blocks
214 std::array<u32, split_count> blocks_ids;
215 blocks_ids[0] = idx_to_refine;
216
217 // generate index for the new blocks (the current index is reused for the first
218 // new block, the others are pushed at the end of the patchdata)
219#pragma unroll
220 for (u32 pid = 0; pid < new_splits; pid++) {
221 blocks_ids[pid + 1] = start_index_push + tid * new_splits + pid;
222 }
223
224 // write coordinates
225
226#pragma unroll
227 for (u32 pid = 0; pid < split_count; pid++) {
228 block_bound_low[blocks_ids[pid]] = block_coords[pid].bmin;
229 block_bound_high[blocks_ids[pid]] = block_coords[pid].bmax;
230 }
231
232 // user lambda to fill the fields
233 uacc.apply_refine(idx_to_refine, cur_block, blocks_ids, block_coords, uacc);
234 });
235 });
236
237 sham::EventList resulting_events{e};
238
239 buf_cell_min.complete_event_state(resulting_events);
240 buf_cell_max.complete_event_state(resulting_events);
241
242 uacc.finalize(resulting_events, pdat);
243
244 refine_flags.complete_event_state(resulting_events);
245 }
246
247 sum_block_count += pdat.get_obj_cnt();
248 new_cell_were_added = new_cell_were_added || refine_flags.get_size() > 0;
249 });
250
251 logger::info_ln("AMRGrid", "process block count =", sum_block_count);
252
253 return new_cell_were_added;
254}
255
256template<class Tvec, class TgridVec>
257template<class UserAcc>
261
262 using namespace shamrock::patch;
263
264 bool cell_were_removed = false;
265
266 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
267 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
268
269 scheduler().for_each_patch_data([&](u64 id_patch, Patch cur_p, PatchDataLayer &pdat) {
270 u32 old_obj_cnt = pdat.get_obj_cnt();
271
272 sham::DeviceBuffer<u32> &derefine_flags = dd_derefine_list.get(id_patch);
273
274 if (derefine_flags.get_size() > 0) {
275
276 // init flag table
277 sham::DeviceBuffer<u32> keep_block_flag(old_obj_cnt, dev_sched);
278 keep_block_flag.fill(1);
279
280 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
281 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
282
283 sham::EventList depends_list;
284 auto block_bound_low = buf_cell_min.get_write_access(depends_list);
285 auto block_bound_high = buf_cell_max.get_write_access(depends_list);
286 UserAcc uacc(depends_list, pdat);
287 auto index_to_deref = derefine_flags.get_read_access(depends_list);
288 auto flag_keep = keep_block_flag.get_write_access(depends_list);
289
290 // edit block content + make flag of blocks to keep
291 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
292 cgh.parallel_for(sycl::range<1>(derefine_flags.get_size()), [=](sycl::item<1> gid) {
293 u32 tid = gid.get_linear_id();
294
295 u32 idx_to_derefine = index_to_deref[gid];
296
297 // compute old block indexes
298 std::array<u32, split_count> old_indexes;
299#pragma unroll
300 for (u32 pid = 0; pid < split_count; pid++) {
301 old_indexes[pid] = idx_to_derefine + pid;
302 }
303
304 // load block coords
305 std::array<BlockCoord, split_count> block_coords;
306#pragma unroll
307 for (u32 pid = 0; pid < split_count; pid++) {
308 block_coords[pid] = BlockCoord{
309 block_bound_low[old_indexes[pid]], block_bound_high[old_indexes[pid]]};
310 }
311
312 // make new block coord
313 BlockCoord merged_block_coord = BlockCoord::get_merge(block_coords);
314
315 // write new coord
316 block_bound_low[idx_to_derefine] = merged_block_coord.bmin;
317 block_bound_high[idx_to_derefine] = merged_block_coord.bmax;
318
319// flag the old blocks for removal
320#pragma unroll
321 for (u32 pid = 1; pid < split_count; pid++) {
322 flag_keep[idx_to_derefine + pid] = 0;
323 }
324
325 // user lambda to fill the fields
326 uacc.apply_derefine(
327 old_indexes, block_coords, idx_to_derefine, merged_block_coord, uacc);
328 });
329 });
330
331 sham::EventList resulting_events{e};
332
333 buf_cell_min.complete_event_state(resulting_events);
334 buf_cell_max.complete_event_state(resulting_events);
335
336 uacc.finalize(resulting_events, pdat);
337
338 keep_block_flag.complete_event_state(resulting_events);
339 derefine_flags.complete_event_state(resulting_events);
340
341 // stream compact the flags
342 auto buf_keep
343 = shamalgs::numeric::stream_compact(dev_sched, keep_block_flag, old_obj_cnt);
344
345 shamlog_debug_ln(
346 "AMR Grid",
347 "patch",
348 id_patch,
349 "derefine block count ",
350 old_obj_cnt,
351 "->",
352 buf_keep.get_size());
353
354 if (buf_keep.get_size() == 0) {
355 throw std::runtime_error("buf keep must contain something at this point");
356 }
357
358 // remap pdat according to stream compact
359 pdat.index_remap_resize(buf_keep, buf_keep.get_size());
360
361 cell_were_removed = cell_were_removed || derefine_flags.get_size() > 0;
362 }
363 });
364
365 return cell_were_removed;
366}
367
368template<class Tvec, class TgridVec>
369template<class UserAccCrit, class UserAccSplit, class UserAccMerge>
372
373 // Ensure that the blocks are sorted before refinement
374 AMRSortBlocks block_sorter(context, solver_config, storage);
375 block_sorter.reorder_amr_blocks();
376
377 // get refine and derefine list
380
381 gen_refine_block_changes<UserAccCrit>(dd_refine_list, dd_derefine_list);
382
384 // Note that this only add new blocks at the end of the patchdata
385 internal_refine_grid<UserAccSplit>(std::move(dd_refine_list));
386
388 // Note that this will perform the merge then remove the old blocks
389 // This is ok to call straight after the refine without edditing the index list in derefine_list
390 // since no permutations were applied in internal_refine_grid and no cells can be both refined
391 // and derefined in the same pass
392 internal_derefine_grid<UserAccMerge>(std::move(dd_derefine_list));
393}
394
395template<class Tvec, class TgridVec>
398
399 class RefineCritBlock {
400 public:
401 const TgridVec *block_low_bound;
402 const TgridVec *block_high_bound;
403 const Tscal *block_density_field;
404
405 Tscal one_over_Nside = 1. / AMRBlock::Nside;
406
407 Tscal dxfact;
408 Tscal wanted_mass;
409
410 RefineCritBlock(
411 sham::EventList &depends_list,
412 u64 id_patch,
413 shamrock::patch::Patch p,
414 shamrock::patch::PatchDataLayer &pdat,
415 Tscal dxfact,
416 Tscal wanted_mass)
417 : dxfact(dxfact), wanted_mass(wanted_mass) {
418
419 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
420 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
421 block_density_field = pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
422 .get_buf()
423 .get_read_access(depends_list);
424 }
425
426 void finalize(
427 sham::EventList &resulting_events,
428 u64 id_patch,
429 shamrock::patch::Patch p,
430 shamrock::patch::PatchDataLayer &pdat,
431 Tscal dxfact,
432 Tscal wanted_mass) {
433
434 sham::DeviceBuffer<i64_3> &buf_cell_low_bound = pdat.get_field<i64_3>(0).get_buf();
435 sham::DeviceBuffer<i64_3> &buf_cell_high_bound = pdat.get_field<i64_3>(1).get_buf();
436
437 buf_cell_low_bound.complete_event_state(resulting_events);
438 buf_cell_high_bound.complete_event_state(resulting_events);
439 pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
440 .get_buf()
441 .complete_event_state(resulting_events);
442 }
443
444 void refine_criterion(
445 u32 block_id, RefineCritBlock acc, bool &should_refine, bool &should_derefine) const {
446
447 TgridVec low_bound = acc.block_low_bound[block_id];
448 TgridVec high_bound = acc.block_high_bound[block_id];
449
450 Tvec lower_flt = low_bound.template convert<Tscal>() * dxfact;
451 Tvec upper_flt = high_bound.template convert<Tscal>() * dxfact;
452
453 Tvec block_cell_size = (upper_flt - lower_flt) * one_over_Nside;
454
455 Tscal sum_mass = 0;
456 for (u32 i = 0; i < AMRBlock::block_size; i++) {
457 sum_mass += acc.block_density_field[i + block_id * AMRBlock::block_size];
458 }
459 sum_mass *= block_cell_size.x() * block_cell_size.y() * block_cell_size.z();
460
461 if (sum_mass > wanted_mass * 8) {
462 should_refine = true;
463 should_derefine = false;
464 } else if (sum_mass < wanted_mass) {
465 should_refine = false;
466 should_derefine = true;
467 } else {
468 should_refine = false;
469 should_derefine = false;
470 }
471
472 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
473 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
474 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
475 }
476 };
477
478 class RefineCellAccessor {
479 public:
480 f64 *rho;
481 f64_3 *rho_vel;
482 f64 *rhoE;
483
484 RefineCellAccessor(sham::EventList &depends_list, shamrock::patch::PatchDataLayer &pdat) {
485
486 rho = pdat.get_field<f64>(2).get_buf().get_write_access(depends_list);
487 rho_vel = pdat.get_field<f64_3>(3).get_buf().get_write_access(depends_list);
488 rhoE = pdat.get_field<f64>(4).get_buf().get_write_access(depends_list);
489 }
490
491 void finalize(sham::EventList &resulting_events, shamrock::patch::PatchDataLayer &pdat) {
492 pdat.get_field<f64>(2).get_buf().complete_event_state(resulting_events);
493 pdat.get_field<f64_3>(3).get_buf().complete_event_state(resulting_events);
494 pdat.get_field<f64>(4).get_buf().complete_event_state(resulting_events);
495 }
496
497 void apply_refine(
498 u32 cur_idx,
499 BlockCoord cur_coords,
500 std::array<u32, 8> new_blocks,
501 std::array<BlockCoord, 8> new_block_coords,
502 RefineCellAccessor acc) const {
503
504 auto get_coord_ref = [](u32 i) -> std::array<u32, dim> {
505 constexpr u32 NsideBlockPow = 1;
506 constexpr u32 Nside = 1U << NsideBlockPow;
507
508 if constexpr (dim == 3) {
509 const u32 tmp = i >> NsideBlockPow;
510 return {i % Nside, (tmp) % Nside, (tmp) >> NsideBlockPow};
511 }
512 };
513
514 auto get_index_block = [](std::array<u32, dim> coord) -> u32 {
515 constexpr u32 NsideBlockPow = 1;
516 constexpr u32 Nside = 1U << NsideBlockPow;
517
518 if constexpr (dim == 3) {
519 return coord[0] + Nside * coord[1] + Nside * Nside * coord[2];
520 }
521 };
522
523 auto get_gid_write = [&](std::array<u32, dim> &glid) -> u32 {
524 std::array<u32, dim> bid
525 = {glid[0] >> AMRBlock::NsideBlockPow,
526 glid[1] >> AMRBlock::NsideBlockPow,
527 glid[2] >> AMRBlock::NsideBlockPow};
528
529 // logger::raw_ln(glid,bid);
530 return new_blocks[get_index_block(bid)] * AMRBlock::block_size
531 + AMRBlock::get_index(
532 {glid[0] % AMRBlock::Nside,
533 glid[1] % AMRBlock::Nside,
534 glid[2] % AMRBlock::Nside});
535 };
536
537 std::array<f64, AMRBlock::block_size> old_rho_block;
538 std::array<f64_3, AMRBlock::block_size> old_rho_vel_block;
539 std::array<f64, AMRBlock::block_size> old_rhoE_block;
540
541 // save old block
542 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
543
544 auto [lx, ly, lz] = get_coord_ref(loc_id);
545 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
546 old_rho_block[loc_id] = acc.rho[old_cell_idx];
547 old_rho_vel_block[loc_id] = acc.rho_vel[old_cell_idx];
548 old_rhoE_block[loc_id] = acc.rhoE[old_cell_idx];
549 }
550
551 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
552
553 auto [lx, ly, lz] = get_coord_ref(loc_id);
554 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
555
556 Tscal rho_block = old_rho_block[loc_id];
557 Tvec rho_vel_block = old_rho_vel_block[loc_id];
558 Tscal rhoE_block = old_rhoE_block[loc_id];
559 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
560
561 auto [sx, sy, sz] = get_coord_ref(subdiv_lid);
562
563 std::array<u32, 3> glid = {lx * 2 + sx, ly * 2 + sy, lz * 2 + sz};
564
565 u32 new_cell_idx = get_gid_write(glid);
566 /*
567 if (1627 == cur_idx) {
568 logger::raw_ln(
569 cur_idx,
570 "set cell ",
571 new_cell_idx,
572 " from cell",
573 old_cell_idx,
574 "old",
575 rho_block,
576 rho_vel_block,
577 rhoE_block);
578 }
579 */
580 acc.rho[new_cell_idx] = rho_block;
581 acc.rho_vel[new_cell_idx] = rho_vel_block;
582 acc.rhoE[new_cell_idx] = rhoE_block;
583 }
584 }
585 }
586
587 void apply_derefine(
588 std::array<u32, 8> old_blocks,
589 std::array<BlockCoord, 8> old_coords,
590 u32 new_cell,
591 BlockCoord new_coord,
592
593 RefineCellAccessor acc) const {
594
595 std::array<f64, AMRBlock::block_size> rho_block;
596 std::array<f64_3, AMRBlock::block_size> rho_vel_block;
597 std::array<f64, AMRBlock::block_size> rhoE_block;
598
599 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
600 rho_block[cell_id] = {};
601 rho_vel_block[cell_id] = {};
602 rhoE_block[cell_id] = {};
603 }
604
605 for (u32 pid = 0; pid < 8; pid++) {
606 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
607 rho_block[cell_id] += acc.rho[old_blocks[pid] * AMRBlock::block_size + cell_id];
608 rho_vel_block[cell_id]
609 += acc.rho_vel[old_blocks[pid] * AMRBlock::block_size + cell_id];
610 rhoE_block[cell_id]
611 += acc.rhoE[old_blocks[pid] * AMRBlock::block_size + cell_id];
612 }
613 }
614
615 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
616 rho_block[cell_id] /= 8;
617 rho_vel_block[cell_id] /= 8;
618 rhoE_block[cell_id] /= 8;
619 }
620
621 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
622 u32 newcell_idx = new_cell * AMRBlock::block_size + cell_id;
623 acc.rho[newcell_idx] = rho_block[cell_id];
624 acc.rho_vel[newcell_idx] = rho_vel_block[cell_id];
625 acc.rhoE[newcell_idx] = rhoE_block[cell_id];
626 }
627 }
628 };
629
630 using AMRmode_None = typename AMRMode<Tvec, TgridVec>::None;
631 using AMRmode_DensityBased = typename AMRMode<Tvec, TgridVec>::DensityBased;
632
633 bool has_cell_order_changed = false;
634
635 if (AMRmode_None *cfg = std::get_if<AMRmode_None>(&solver_config.amr_mode.config)) {
636 // no refinment here turn around there is nothing to see
637 } else if (
638 AMRmode_DensityBased *cfg
639 = std::get_if<AMRmode_DensityBased>(&solver_config.amr_mode.config)) {
640 Tscal dxfact(solver_config.grid_coord_to_pos_fact);
641
642 // get refine and derefine list
645
646 gen_refine_block_changes_old<RefineCritBlock>(
647 refine_list, derefine_list, dxfact, cfg->crit_mass);
648
650 // Note that this only add new blocks at the end of the patchdata
651 bool change_refine = internal_refine_grid_old<RefineCellAccessor>(std::move(refine_list));
652
654 // Note that this will perform the merge then remove the old blocks
655 // This is ok to call straight after the refine without edditing the index list in
656 // derefine_list since no permutations were applied in internal_refine_grid and no cells can
657 // be both refined and derefined in the same pass
658 bool change_derefine
659 = internal_derefine_grid_old<RefineCellAccessor>(std::move(derefine_list));
660
661 has_cell_order_changed = has_cell_order_changed || (change_refine || change_derefine);
662 }
663
664 if (has_cell_order_changed) {
665 // Ensure that the blocks are sorted before refinement
666 AMRSortBlocks block_sorter(context, solver_config, storage);
667 block_sorter.reorder_amr_blocks();
668 }
669}
670
671template<class Tvec, class TgridVec>
672template<class UserAcc, class... T>
677 T &&...args) {
678
679 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
680 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
681
682 scheduler().for_each_patchdata_nonempty([&](Patch cur_p, PatchDataLayer &pdat) {
683 u64 id_patch = cur_p.id_patch;
684
685 // create the refine and derefine flags buffers
686 u32 obj_cnt = pdat.get_obj_cnt();
687
688 sham::DeviceBuffer<u32> refine_flags(obj_cnt, dev_sched);
689 sham::DeviceBuffer<u32> derefine_flags(obj_cnt, dev_sched);
690
691 {
692 sham::EventList depends_list;
693
694 UserAcc uacc(depends_list, storage, id_patch, cur_p, pdat, args...);
695
696 auto refine_acc = refine_flags.get_write_access(depends_list);
697 auto derefine_acc = derefine_flags.get_write_access(depends_list);
698
699 // fill in the flags
700 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
701 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
702 bool flag_refine = false;
703 bool flag_derefine = false;
704 uacc.refine_criterion_new(
705 gid.get_linear_id(), uacc, flag_refine, flag_derefine);
706
707 // This is just a safe guard to avoid this nonsensicall case
708 if (flag_refine && flag_derefine) {
709 flag_derefine = false;
710 }
711
712 refine_acc[gid] = (flag_refine) ? 1 : 0;
713 derefine_acc[gid] = (flag_derefine) ? 1 : 0;
714 });
715 });
716
717 sham::EventList resulting_events;
718 resulting_events.add_event(e);
719
720 refine_flags.complete_event_state(resulting_events);
721 derefine_flags.complete_event_state(resulting_events);
722
723 uacc.finalize_new(resulting_events, storage, id_patch, cur_p, pdat, args...);
724 }
725
726 dd_refine_flags.add_obj(id_patch, std::move(refine_flags));
727 dd_derefine_flags.add_obj(id_patch, std::move(derefine_flags));
728 });
729}
730
737template<class Tvec, class TgridVec>
741
742 scheduler().for_each_patchdata_nonempty([&](Patch cur_p, PatchDataLayer &pdat) {
743 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
744 u64 id_patch = cur_p.id_patch;
745 sham::DeviceBuffer<u32> &patch_refine_flags = dd_refine_flags.get(id_patch);
746 u32 obj_cnt = pdat.get_obj_cnt();
747
748 // blocks graph in each direction for the current patch
749 AMRGraph &block_graph_neighs_xp = shambase::get_check_ref(storage.block_graph_edge)
750 .get_refs_dir(Direction_::xp)
751 .get(id_patch);
752 AMRGraph &block_graph_neighs_xm = shambase::get_check_ref(storage.block_graph_edge)
753 .get_refs_dir(Direction_::xm)
754 .get(id_patch);
755 AMRGraph &block_graph_neighs_yp = shambase::get_check_ref(storage.block_graph_edge)
756 .get_refs_dir(Direction_::yp)
757 .get(id_patch);
758 AMRGraph &block_graph_neighs_ym = shambase::get_check_ref(storage.block_graph_edge)
759 .get_refs_dir(Direction_::ym)
760 .get(id_patch);
761 AMRGraph &block_graph_neighs_zp = shambase::get_check_ref(storage.block_graph_edge)
762 .get_refs_dir(Direction_::zp)
763 .get(id_patch);
764 AMRGraph &block_graph_neighs_zm = shambase::get_check_ref(storage.block_graph_edge)
765 .get_refs_dir(Direction_::zm)
766 .get(id_patch);
767 // get levels in the current patch
768 sham::DeviceBuffer<TgridUint> &buf_amr_block_levels
769 = shambase::get_check_ref(storage.amr_block_levels).get_buf(id_patch);
770
771 std::shared_ptr<sham::DeviceScheduler> dev_sched
772 = shamsys::instance::get_compute_scheduler_ptr();
773
774 sham::DeviceBuffer<u32> changed_buf(1, dev_sched);
775
776 for (u32 pass = 0; pass < 100; pass++) {
777
778 changed_buf.set_val_at_idx(0, 0);
779
780 sham::EventList depend_list;
781 AMRGraphLinkiterator block_graph_xp
782 = block_graph_neighs_xp.get_read_access(depend_list);
783 AMRGraphLinkiterator block_graph_xm
784 = block_graph_neighs_xm.get_read_access(depend_list);
785 AMRGraphLinkiterator block_graph_yp
786 = block_graph_neighs_yp.get_read_access(depend_list);
787 AMRGraphLinkiterator block_graph_ym
788 = block_graph_neighs_ym.get_read_access(depend_list);
789 AMRGraphLinkiterator block_graph_zp
790 = block_graph_neighs_zp.get_read_access(depend_list);
791 AMRGraphLinkiterator block_graph_zm
792 = block_graph_neighs_zm.get_read_access(depend_list);
793
794 auto acc_amr_levels = buf_amr_block_levels.get_read_access(depend_list);
795 auto acc_changed = changed_buf.get_write_access(depend_list);
796
797 auto acc_ref_flags = patch_refine_flags.get_write_access(depend_list);
798
799 auto e = q.submit(depend_list, [&](sycl::handler &cgh) {
800 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
801 u32 block_id = gid.get_linear_id();
802
803 u32 cur_ref_flag = acc_ref_flags[block_id];
804
805 if (!cur_ref_flag)
806 return;
807
808 auto cur_block_level = acc_amr_levels[block_id];
809
810 auto check_2To1_ref = [&](u32 nid) {
811 if (nid >= obj_cnt)
812 return;
813
814 // get refinement flag and amr level of the neighborh block
815 u32 neigh_ref_flag = acc_ref_flags[nid];
816 auto neigh_block_level = acc_amr_levels[nid];
817
818 auto cur_future = cur_block_level + (cur_ref_flag ? 1 : 0);
819
820 auto neigh_future = neigh_block_level + (neigh_ref_flag ? 1 : 0);
821
822 if (cur_ref_flag && (cur_future > neigh_future + 1)) {
823
824 if (!neigh_ref_flag) {
825 sycl::atomic_ref<
826 u32,
827 sycl::memory_order::relaxed,
828 sycl::memory_scope::system>
829 atomic_neigh_flag(acc_ref_flags[nid]);
830 atomic_neigh_flag.exchange(1);
831
832 sycl::atomic_ref<
833 u32,
834 sycl::memory_order::relaxed,
835 sycl::memory_scope::system>
836 atomic_changed(acc_changed[0]);
837 atomic_changed.exchange(1);
838 }
839 }
840 };
841
842 block_graph_xp.for_each_object_link(block_id, check_2To1_ref);
843 block_graph_xm.for_each_object_link(block_id, check_2To1_ref);
844 block_graph_yp.for_each_object_link(block_id, check_2To1_ref);
845 block_graph_ym.for_each_object_link(block_id, check_2To1_ref);
846 block_graph_zp.for_each_object_link(block_id, check_2To1_ref);
847 block_graph_zm.for_each_object_link(block_id, check_2To1_ref);
848 });
849 });
850 block_graph_neighs_xp.complete_event_state(e);
851 block_graph_neighs_xm.complete_event_state(e);
852 block_graph_neighs_yp.complete_event_state(e);
853 block_graph_neighs_ym.complete_event_state(e);
854 block_graph_neighs_zp.complete_event_state(e);
855 block_graph_neighs_zm.complete_event_state(e);
856 buf_amr_block_levels.complete_event_state(e);
857
858 changed_buf.complete_event_state(e);
859
860 patch_refine_flags.complete_event_state(e);
861
862 e.wait();
863
864 if (changed_buf.get_val_at_idx(0) == 0) {
865 logger::raw_ln("Refinement 2:1 balance converged in ", pass + 1, " sweeps");
866 break;
867 }
868 }
869
871 // refinement
873
874 // perform stream compactions on the refinement flags
875 auto dev_buf_ref
876 = shamalgs::numeric::stream_compact(dev_sched, patch_refine_flags, obj_cnt);
877
878 shamlog_debug_ln(
879 "AMRGrid", "patch ", id_patch, dev_buf_ref.get_size(), "marked for refinement + 2:1");
880 });
881}
882
890
891template<class Tvec, class TgridVec>
896
897 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
898
899 scheduler().for_each_patchdata_nonempty([&](Patch cur_p, PatchDataLayer &pdat) {
900 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
901 u64 id_patch = cur_p.id_patch;
902
903 sham::DeviceBuffer<u32> &patch_derefine_flag = dd_derefine_flags.get(id_patch);
904 sham::DeviceBuffer<u32> &patch_refine_flag = dd_refine_flags.get(id_patch);
905
906 u32 obj_cnt = pdat.get_obj_cnt();
907
908 // blocks graph in each direction for the current patch
909 AMRGraph &block_graph_neighs_xp = shambase::get_check_ref(storage.block_graph_edge)
910 .get_refs_dir(Direction_::xp)
911 .get(id_patch);
912 AMRGraph &block_graph_neighs_xm = shambase::get_check_ref(storage.block_graph_edge)
913 .get_refs_dir(Direction_::xm)
914 .get(id_patch);
915 AMRGraph &block_graph_neighs_yp = shambase::get_check_ref(storage.block_graph_edge)
916 .get_refs_dir(Direction_::yp)
917 .get(id_patch);
918 AMRGraph &block_graph_neighs_ym = shambase::get_check_ref(storage.block_graph_edge)
919 .get_refs_dir(Direction_::ym)
920 .get(id_patch);
921 AMRGraph &block_graph_neighs_zp = shambase::get_check_ref(storage.block_graph_edge)
922 .get_refs_dir(Direction_::zp)
923 .get(id_patch);
924 AMRGraph &block_graph_neighs_zm = shambase::get_check_ref(storage.block_graph_edge)
925 .get_refs_dir(Direction_::zm)
926 .get(id_patch);
927 // get the current buffer of block levels in the current patch
928 sham::DeviceBuffer<TgridUint> &buf_amr_block_levels
929 = shambase::get_check_ref(storage.amr_block_levels).get_buf(id_patch);
930
933
934 auto dev_buf_deref_0
935 = shamalgs::numeric::stream_compact(dev_sched, patch_derefine_flag, obj_cnt);
936
938 " Count block's flag for derefinement [No geometry validity check and no 2:1 check] \t "
939 ": ",
940 dev_buf_deref_0.get_size(),
941 "\n");
944 // keep derefine flags on only if the eight cells want to merge and if they can
945 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
946 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
947
948 sham::EventList depends_list;
949 auto acc_min = buf_cell_min.get_read_access(depends_list);
950 auto acc_max = buf_cell_max.get_read_access(depends_list);
951 auto acc_amr_levels = buf_amr_block_levels.get_read_access(depends_list);
952
953 auto acc_merge_flag = patch_derefine_flag.get_write_access(depends_list);
954 auto acc_refine_flag = patch_refine_flag.get_read_access(depends_list);
955
956 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
957 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
958 u32 id = gid.get_linear_id();
959
960 std::array<BlockCoord, split_count> blocks;
961 bool do_merge = true;
962 bool all_same_level = true;
963
964 // This avoid the case where we are in the last block of the buffer to
965 // avoid the out-of-bound read
966 if (id + split_count <= obj_cnt) {
967 bool all_want_to_merge = true;
968
969 auto get_coord = [](u32 i) -> std::array<u32, dim> {
970 constexpr u32 NsideBlockPow = 1;
971 constexpr u32 Nside = 1U << NsideBlockPow;
972 constexpr u32 side_size = Nside;
973 constexpr u32 block_size = shambase::pow_constexpr<dim>(Nside);
974
975 if constexpr (dim == 3) {
976 const u32 tmp = i >> NsideBlockPow;
977 // This line is why derefinement never happens
978 // return {i % Nside, (tmp) % Nside, (tmp ) >> NsideBlockPow};
979 return {(tmp) >> NsideBlockPow, (tmp) % Nside, i % Nside};
980 }
981 };
982
983 auto get_split
984 = [=](BlockCoord target_block) -> std::array<BlockCoord, split_count> {
985 std::array<BlockCoord, split_count> ret;
986 auto bmin = target_block.bmin;
987 auto bmax = target_block.bmax;
988 auto split = bmin + (bmax - bmin) / 2;
989 std::array<TgridVec, 3> szs = {bmin, split, bmax};
990 for (u32 i = 0; i < split_count; i++) {
991 auto [lx, ly, lz] = get_coord(i);
992
993 ret[i].bmin = TgridVec{szs[lx].x(), szs[ly].y(), szs[lz].z()};
994 ret[i].bmax
995 = TgridVec{szs[lx + 1].x(), szs[ly + 1].y(), szs[lz + 1].z()};
996 }
997
998 return ret;
999 };
1000
1001 for (u32 b_lid = 0; b_lid < split_count; b_lid++) {
1002 blocks[b_lid] = BlockCoord{acc_min[id + b_lid], acc_max[id + b_lid]};
1003 all_want_to_merge = all_want_to_merge && acc_merge_flag[id + b_lid];
1004 all_same_level
1005 = all_same_level && (acc_amr_levels[id] == acc_amr_levels[id + b_lid]);
1006 }
1007
1008 BlockCoord merged = BlockCoord::get_merge(blocks);
1009 std::array<BlockCoord, split_count> splitted = get_split(merged);
1010 for (u32 lid = 0; lid < split_count; lid++) {
1011 do_merge = do_merge && sham::equals(blocks[lid].bmin, splitted[lid].bmin)
1012 && sham::equals(blocks[lid].bmax, splitted[lid].bmax);
1013 }
1014
1015 do_merge = do_merge && all_want_to_merge && all_same_level;
1016 if (acc_refine_flag[id] && do_merge) {
1017 do_merge = false;
1018 }
1019
1020 } else {
1021 do_merge = false;
1022 }
1023 acc_merge_flag[id] = do_merge;
1024 });
1025 });
1026 buf_cell_min.complete_event_state(e);
1027 buf_cell_max.complete_event_state(e);
1028 buf_amr_block_levels.complete_event_state(e);
1029 patch_derefine_flag.complete_event_state(e);
1030 patch_refine_flag.complete_event_state(e);
1031
1034 auto buf_derefine_1
1035 = shamalgs::numeric::stream_compact(dev_sched, patch_derefine_flag, obj_cnt);
1037 " Count block's flag for derefinement [After geometry validity check and before 2:1 "
1038 "check] "
1039 "\t : ",
1040 buf_derefine_1.get_size(),
1041 "\n");
1044
1045 // ////////////////////////////////////////////////////////////////////////////////////
1046 // // // enforce 2:1 at parent level
1047 // ///////////////////////////////////////////////////////////////////////////////////
1048
1049 std::shared_ptr<sham::DeviceScheduler> dev_sched
1050 = shamsys::instance::get_compute_scheduler_ptr();
1051
1052 sham::DeviceBuffer<u32> changed_buf(1, dev_sched);
1053
1054 // copy old deref buffer to avoid race condition
1055 sham::DeviceBuffer<u32> patch_derefine_flag_old(obj_cnt, dev_sched);
1056 patch_derefine_flag.copy_range(0, obj_cnt, patch_derefine_flag_old);
1057
1058 //
1059 sham::DeviceBuffer<u32> patch_derefine_flag_new(obj_cnt, dev_sched);
1060
1061 for (int it = 0; it < 100; it++) {
1062 changed_buf.set_val_at_idx(0, 0);
1063
1064 sham::EventList depend_list;
1065
1066 AMRGraphLinkiterator block_graph_xp
1067 = block_graph_neighs_xp.get_read_access(depend_list);
1068
1069 AMRGraphLinkiterator block_graph_xm
1070 = block_graph_neighs_xm.get_read_access(depend_list);
1071 AMRGraphLinkiterator block_graph_yp
1072 = block_graph_neighs_yp.get_read_access(depend_list);
1073 AMRGraphLinkiterator block_graph_ym
1074 = block_graph_neighs_ym.get_read_access(depend_list);
1075 AMRGraphLinkiterator block_graph_zp
1076 = block_graph_neighs_zp.get_read_access(depend_list);
1077 AMRGraphLinkiterator block_graph_zm
1078 = block_graph_neighs_zm.get_read_access(depend_list);
1079
1080 auto acc_amr_levels = buf_amr_block_levels.get_read_access(depend_list);
1081 auto acc_ref_flag = patch_refine_flag.get_read_access(depend_list);
1082 auto acc_changed = changed_buf.get_write_access(depend_list);
1083
1084 auto acc_deref_old = patch_derefine_flag_old.get_read_access(depend_list);
1085 auto acc_deref_new = patch_derefine_flag_new.get_write_access(depend_list);
1086
1087 auto e_2to1 = q.submit(depend_list, [&](sycl::handler &cgh) {
1088 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
1089 auto lid = gid.get_linear_id();
1090
1091 auto old_flag = acc_deref_old[lid];
1092 auto new_flag = old_flag;
1093
1094 auto check_2To1_der = [&](u32 nid) {
1095 if (nid < obj_cnt)
1096
1097 {
1098
1099 auto neigh_future = acc_amr_levels[nid] + (acc_ref_flag[nid] ? 1 : 0)
1100 - (acc_deref_old[nid] ? 1 : 0);
1101
1102 auto my_future = acc_amr_levels[lid] - 1;
1103
1104 if (neigh_future > my_future + 1) {
1105 new_flag = 0;
1106 }
1107 }
1108 };
1109
1110 if (old_flag) {
1111
1112 for (u32 i = 0; i < AMRBlock::block_size; i++) {
1113 block_graph_xp.for_each_object_link((lid + i), check_2To1_der);
1114 block_graph_xm.for_each_object_link((lid + i), check_2To1_der);
1115 block_graph_yp.for_each_object_link((lid + i), check_2To1_der);
1116 block_graph_ym.for_each_object_link((lid + i), check_2To1_der);
1117 block_graph_zp.for_each_object_link((lid + i), check_2To1_der);
1118 block_graph_zm.for_each_object_link((lid + i), check_2To1_der);
1119 }
1120
1121 if (old_flag != new_flag) {
1122 sycl::atomic_ref<
1123 u32,
1124 sycl::memory_order::relaxed,
1125 sycl::memory_scope::system>
1126 atomic_changed(acc_changed[0]);
1127 atomic_changed.exchange(1);
1128 }
1129 }
1130
1131 acc_deref_new[lid] = new_flag;
1132 });
1133 });
1134 block_graph_neighs_xp.complete_event_state(e_2to1);
1135 block_graph_neighs_xm.complete_event_state(e_2to1);
1136 block_graph_neighs_yp.complete_event_state(e_2to1);
1137 block_graph_neighs_ym.complete_event_state(e_2to1);
1138 block_graph_neighs_zp.complete_event_state(e_2to1);
1139 block_graph_neighs_zm.complete_event_state(e_2to1);
1140 buf_amr_block_levels.complete_event_state(e_2to1);
1141 patch_refine_flag.complete_event_state(e_2to1);
1142 changed_buf.complete_event_state(e_2to1);
1143
1144 patch_derefine_flag_old.complete_event_state(e_2to1);
1145 patch_derefine_flag_new.complete_event_state(e_2to1);
1146 e_2to1.wait();
1147
1148 std::swap(patch_derefine_flag_old, patch_derefine_flag_new);
1149
1150 if (changed_buf.get_val_at_idx(0) == 0) {
1151 logger:
1153 "Derefinement 2:1 balance converge in \t ", it + 1, "\t sweeps \n\n");
1154 break;
1155 }
1156 }
1157
1158 // copy back to ..
1159 patch_derefine_flag_old.copy_range(0, obj_cnt, patch_derefine_flag);
1160
1161 // ////////////////////////////////////////////////////////////////////////////////
1162 // // derefinement
1163 // ////////////////////////////////////////////////////////////////////////////////
1164 // perform stream compactions on the derefinement flags
1165 auto buf_derefine
1166 = shamalgs::numeric::stream_compact(dev_sched, patch_derefine_flag, obj_cnt);
1167
1169 " Count block's flag for derefinement [After geometry validity check and after 2:1 "
1170 "check] \t : ",
1171 buf_derefine.get_size(),
1172 "\n");
1173
1174 shamlog_debug_ln(
1175 "AMRGrid", "patch ", id_patch, buf_derefine.get_size(), "marked for derefinement ");
1176 });
1177}
1178
1179template<class Tvec, class TgridVec>
1180template<class UserAcc>
1183
1184 u64 sum_block_count = 0;
1185
1186 bool new_cell_were_added = false;
1187 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
1188 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
1189
1190 scheduler().for_each_patch_data([&](u64 id_patch, Patch cur_p, PatchDataLayer &pdat) {
1191 u32 old_obj_cnt = pdat.get_obj_cnt();
1192 auto stream_compaction_result = shamalgs::numeric::stream_compact(
1193 dev_sched, dd_refine_flags.get(id_patch), old_obj_cnt);
1194
1195 if (stream_compaction_result.get_size() > 0) {
1196 // alloc memory for the new blocks to be created
1197
1199 "Will refine \t", stream_compaction_result.get_size(), " \t blocks \n\n");
1200
1201 pdat.expand(static_cast<u32>(stream_compaction_result.get_size()) * (split_count - 1));
1202 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
1203 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
1204 sham::EventList depends_list;
1205
1206 auto block_bound_low = buf_cell_min.get_write_access(depends_list);
1207 auto block_bound_high = buf_cell_max.get_write_access(depends_list);
1208 UserAcc uacc(depends_list, storage, id_patch, pdat);
1209 auto index_to_ref = stream_compaction_result.get_read_access(depends_list);
1210
1211 // Refine the block (set the positions) and fill the corresponding fields
1212 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
1213 u32 start_index_push = old_obj_cnt;
1214
1215 constexpr u32 new_splits = split_count - 1;
1216
1217 cgh.parallel_for(
1218 sycl::range<1>(stream_compaction_result.get_size()), [=](sycl::item<1> gid) {
1219 u32 tid = gid.get_linear_id();
1220
1221 u32 idx_to_refine = index_to_ref[gid];
1222
1223 // gen splits coordinates
1224 BlockCoord cur_block{
1225 block_bound_low[idx_to_refine], block_bound_high[idx_to_refine]};
1226
1227 std::array<BlockCoord, split_count> block_coords
1228 = BlockCoord::get_split(cur_block.bmin, cur_block.bmax);
1229
1230 // generate index for the refined blocks
1231 std::array<u32, split_count> blocks_ids;
1232 blocks_ids[0] = idx_to_refine;
1233
1234 // generate index for the new blocks (the current index is reused for the first
1235 // new block, the others are pushed at the end of the patchdata)
1236#pragma unroll
1237 for (u32 pid = 0; pid < new_splits; pid++) {
1238 blocks_ids[pid + 1] = start_index_push + tid * new_splits + pid;
1239 }
1240
1241 // write coordinates
1242
1243#pragma unroll
1244 for (u32 pid = 0; pid < split_count; pid++) {
1245 block_bound_low[blocks_ids[pid]] = block_coords[pid].bmin;
1246 block_bound_high[blocks_ids[pid]] = block_coords[pid].bmax;
1247 }
1248
1249 // user lambda to fill the fields
1250 uacc.apply_refine_new(
1251 idx_to_refine, cur_block, blocks_ids, block_coords, uacc);
1252 });
1253 });
1254
1255 sham::EventList resulting_events;
1256
1257 buf_cell_min.complete_event_state(resulting_events);
1258 buf_cell_max.complete_event_state(resulting_events);
1259 stream_compaction_result.complete_event_state(e);
1260 uacc.finalize_new(resulting_events, storage, id_patch, pdat);
1261 }
1262
1263 shamlog_debug_ln("AMRGrid", "patch ", id_patch, "new block count = ", pdat.get_obj_cnt());
1264 sum_block_count += pdat.get_obj_cnt();
1265 new_cell_were_added = new_cell_were_added || (stream_compaction_result.get_size() > 0);
1266 });
1267
1268 logger::info_ln("AMRGrid", "process block count =", sum_block_count);
1269
1270 return new_cell_were_added;
1271}
1272
1273template<class Tvec, class TgridVec>
1274template<class UserAcc>
1278
1279 using namespace shamrock::patch;
1280
1281 bool cell_were_removed = false;
1282
1283 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
1284 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
1285
1286 scheduler().for_each_patch_data([&](u64 id_patch, Patch cur_p, PatchDataLayer &pdat) {
1287 u32 old_obj_cnt = pdat.get_obj_cnt();
1288 u32 old_obj_cnt_before_refinement = dd_derefine_flags.get(id_patch).get_size();
1289 auto stream_compact_results = shamalgs::numeric::stream_compact(
1290 dev_sched, dd_derefine_flags.get(id_patch), old_obj_cnt_before_refinement);
1291 if (stream_compact_results.get_size() > 0) {
1292 // init flag table
1293 sycl::buffer<u32> keep_block_flag
1294 = shamalgs::algorithm::gen_buffer_device(q.q, old_obj_cnt, [](u32 i) -> u32 {
1295 return 1;
1296 });
1297
1298 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
1299 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
1300 sham::EventList depends_list;
1301 auto block_bound_low = buf_cell_min.get_write_access(depends_list);
1302 auto block_bound_high = buf_cell_max.get_write_access(depends_list);
1303 UserAcc uacc(depends_list, storage, id_patch, pdat);
1304 auto index_to_deref = stream_compact_results.get_read_access(depends_list);
1305
1306 // edit block content + make flag of blocks to keep
1307 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
1308 sycl::accessor flag_keep{keep_block_flag, cgh, sycl::read_write};
1309 cgh.parallel_for(
1310 sycl::range<1>(stream_compact_results.get_size()), [=](sycl::item<1> gid) {
1311 u32 tid = gid.get_linear_id();
1312
1313 u32 idx_to_derefine = index_to_deref[gid];
1314
1315 // compute old block indexes
1316 std::array<u32, split_count> old_indexes;
1317#pragma unroll
1318 for (u32 pid = 0; pid < split_count; pid++) {
1319 old_indexes[pid] = idx_to_derefine + pid;
1320 }
1321
1322 // load block coords
1323 std::array<BlockCoord, split_count> block_coords;
1324#pragma unroll
1325 for (u32 pid = 0; pid < split_count; pid++) {
1326 block_coords[pid] = BlockCoord{
1327 block_bound_low[old_indexes[pid]],
1328 block_bound_high[old_indexes[pid]]};
1329 }
1330
1331 // make new block coord
1332 BlockCoord merged_block_coord = BlockCoord::get_merge(block_coords);
1333
1334 // write new coord
1335 block_bound_low[idx_to_derefine] = merged_block_coord.bmin;
1336 block_bound_high[idx_to_derefine] = merged_block_coord.bmax;
1337
1338// flag the old blocks for removal
1339#pragma unroll
1340 for (u32 pid = 1; pid < split_count; pid++) {
1341 flag_keep[idx_to_derefine + pid] = 0;
1342 }
1343
1344 // user lambda to fill the fields
1345
1346 uacc.apply_derefine_new(
1347 old_indexes, block_coords, idx_to_derefine, merged_block_coord, uacc);
1348 });
1349 });
1350
1351 sham::EventList resulting_events;
1352
1353 buf_cell_min.complete_event_state(resulting_events);
1354 buf_cell_max.complete_event_state(resulting_events);
1355 uacc.finalize_new(resulting_events, storage, id_patch, pdat);
1356
1357 stream_compact_results.complete_event_state(resulting_events);
1358
1359 // stream compact the flags (get new block ids map after merged)
1360 auto [opt_buf, len]
1361 = shamalgs::numeric::stream_compact(q.q, keep_block_flag, old_obj_cnt);
1362
1364 "AMR Grid",
1365 "patch",
1366 id_patch,
1367 "derefine block count = ",
1368 old_obj_cnt - len,
1369 "new block count = ",
1370 len);
1371
1372 if (!opt_buf) {
1373 throw std::runtime_error("opt buf must contain something at this point");
1374 }
1375
1376 // remap pdat according to stream compact (for each field in patchdataleyer resize
1377 // according to new block ids map)
1378 pdat.index_remap_resize(*opt_buf, len);
1379
1380 cell_were_removed = cell_were_removed || stream_compact_results.get_size() > 0;
1381 }
1382 });
1383
1384 return cell_were_removed;
1385}
1386
1387template<class Tvec, class TgridVec>
1388template<class UserAccCrit, class UserAccSplit, class UserAccMerge>
1391
1392 // Ensure that the blocks are sorted before refinement
1393 AMRSortBlocks block_sorter(context, solver_config, storage);
1394 block_sorter.reorder_amr_blocks();
1395
1396 // get refine and derefine list
1399
1400 gen_refine_block_changes_new<UserAccCrit>(dd_refine_list, dd_derefine_list);
1401
1403 // Note that this only add new blocks at the end of the patchdata
1404 internal_refine_grid_new<UserAccSplit>(std::move(dd_refine_list));
1405
1407 // Note that this will perform the merge then remove the old blocks
1408 // This is ok to call straight after the refine without edditing the index list in derefine_list
1409 // since no permutations were applied in internal_refine_grid_new and no cells can be both
1410 // refined and derefined in the same pass
1411 internal_derefine_grid_new<UserAccMerge>(std::move(dd_derefine_list));
1412}
1413
1414template<class Tvec, class TgridVec>
1417
1418 class RefineCritBlock {
1419 public:
1420 const TgridVec *block_low_bound;
1421 const TgridVec *block_high_bound;
1422 const Tscal *block_density_field;
1423
1424 Tscal one_over_Nside = 1. / AMRBlock::Nside;
1425
1426 Tscal dxfact;
1427 Tscal wanted_mass;
1428
1429 RefineCritBlock(
1430 sham::EventList &depends_list,
1431 Storage &storage,
1432 u64 id_patch,
1435 Tscal dxfact,
1436 Tscal wanted_mass)
1437 : dxfact(dxfact), wanted_mass(wanted_mass) {
1438
1439 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
1440 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
1441 block_density_field = pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1442 .get_buf()
1443 .get_read_access(depends_list);
1444 }
1445
1446 void finalize_new(
1447 sham::EventList &resulting_events,
1448 Storage &storage,
1449 u64 id_patch,
1452 Tscal dxfact,
1453 Tscal wanted_mass) {
1454
1455 pdat.get_field<TgridVec>(0).get_buf().complete_event_state(resulting_events);
1456 pdat.get_field<TgridVec>(1).get_buf().complete_event_state(resulting_events);
1457 pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1458 .get_buf()
1459 .complete_event_state(resulting_events);
1460 }
1461
1462 void refine_criterion_new(
1463 u32 block_id, RefineCritBlock acc, bool &should_refine, bool &should_derefine) const {
1464
1465 TgridVec low_bound = acc.block_low_bound[block_id];
1466 TgridVec high_bound = acc.block_high_bound[block_id];
1467
1468 Tvec lower_flt = low_bound.template convert<Tscal>() * dxfact;
1469 Tvec upper_flt = high_bound.template convert<Tscal>() * dxfact;
1470
1471 Tvec block_cell_size = (upper_flt - lower_flt) * one_over_Nside;
1472
1473 Tscal sum_mass = 0;
1474 for (u32 i = 0; i < AMRBlock::block_size; i++) {
1475 sum_mass += acc.block_density_field[i + block_id * AMRBlock::block_size];
1476 }
1477 sum_mass *= block_cell_size.x() * block_cell_size.y() * block_cell_size.z();
1478
1479 if (sum_mass > wanted_mass * 8) {
1480 should_refine = true;
1481 should_derefine = false;
1482 } else if (sum_mass < wanted_mass) {
1483 should_refine = false;
1484 should_derefine = true;
1485 } else {
1486 should_refine = false;
1487 should_derefine = false;
1488 }
1489
1490 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
1491 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
1492 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
1493 }
1494 };
1495
1496 class RefineCellAccessor {
1497 public:
1498 f64 *rho;
1499 f64_3 *rho_vel;
1500 f64 *rhoE;
1501 u64 p_id;
1502 // f64* cell_sizes;
1503
1504 // this will be needed for interpolation during refinement
1505 AMRGraphLinkiterator cell_graph_xp;
1506 AMRGraphLinkiterator cell_graph_xm;
1507 AMRGraphLinkiterator cell_graph_yp;
1508 AMRGraphLinkiterator cell_graph_ym;
1509 AMRGraphLinkiterator cell_graph_zp;
1510 AMRGraphLinkiterator cell_graph_zm;
1511
1512 RefineCellAccessor(
1513 sham::EventList &depends_list,
1514 Storage &storage,
1515 u64 &id_patch,
1517 : cell_graph_xp(
1518 shambase::get_check_ref(storage.cell_graph_edge)
1519 .get_refs_dir(Direction::xp)
1520 .get(id_patch)
1521 .get()
1522 .get_read_access(depends_list)),
1523 cell_graph_xm(
1524 shambase::get_check_ref(storage.cell_graph_edge)
1525 .get_refs_dir(Direction::xm)
1526 .get(id_patch)
1527 .get()
1528 .get_read_access(depends_list)),
1529 cell_graph_yp(
1530 shambase::get_check_ref(storage.cell_graph_edge)
1531 .get_refs_dir(Direction::yp)
1532 .get(id_patch)
1533 .get()
1534 .get_read_access(depends_list)),
1535 cell_graph_ym(
1536 shambase::get_check_ref(storage.cell_graph_edge)
1537 .get_refs_dir(Direction::ym)
1538 .get(id_patch)
1539 .get()
1540 .get_read_access(depends_list)),
1541 cell_graph_zp(
1542 shambase::get_check_ref(storage.cell_graph_edge)
1543 .get_refs_dir(Direction::zp)
1544 .get(id_patch)
1545 .get()
1546 .get_read_access(depends_list)),
1547 cell_graph_zm(
1548 shambase::get_check_ref(storage.cell_graph_edge)
1549 .get_refs_dir(Direction::zm)
1550 .get(id_patch)
1551 .get()
1552 .get_read_access(depends_list))
1553
1554 {
1555 p_id = id_patch;
1556 rho = pdat.get_field<f64>(2).get_buf().get_write_access(depends_list);
1557 rho_vel = pdat.get_field<f64_3>(3).get_buf().get_write_access(depends_list);
1558 rhoE = pdat.get_field<f64>(4).get_buf().get_write_access(depends_list);
1559 // cell_sizes = shambase::get_check_ref(storage.block_cell_sizes)
1560 // .get_buf(id_patch)
1561 // .get_write_access(depends_list);
1562 }
1563
1564 void finalize_new(
1565 sham::EventList &resulting_events,
1566 Storage &storage,
1567 u64 &id_patch,
1569 pdat.get_field<f64>(2).get_buf().complete_event_state(resulting_events);
1570 pdat.get_field<f64_3>(3).get_buf().complete_event_state(resulting_events);
1571 pdat.get_field<f64>(4).get_buf().complete_event_state(resulting_events);
1572
1573 shambase::get_check_ref(storage.cell_graph_edge)
1574 .get_refs_dir(Direction_::xp)
1575 .get(id_patch)
1576 .get()
1577 .complete_event_state(resulting_events);
1578 shambase::get_check_ref(storage.cell_graph_edge)
1579 .get_refs_dir(Direction_::xm)
1580 .get(id_patch)
1581 .get()
1582 .complete_event_state(resulting_events);
1583 shambase::get_check_ref(storage.cell_graph_edge)
1584 .get_refs_dir(Direction_::yp)
1585 .get(id_patch)
1586 .get()
1587 .complete_event_state(resulting_events);
1588 shambase::get_check_ref(storage.cell_graph_edge)
1589 .get_refs_dir(Direction_::ym)
1590 .get(id_patch)
1591 .get()
1592 .complete_event_state(resulting_events);
1593 shambase::get_check_ref(storage.cell_graph_edge)
1594 .get_refs_dir(Direction_::zp)
1595 .get(id_patch)
1596 .get()
1597 .complete_event_state(resulting_events);
1598 shambase::get_check_ref(storage.cell_graph_edge)
1599 .get_refs_dir(Direction_::zm)
1600 .get(id_patch)
1601 .get()
1602 .complete_event_state(resulting_events);
1603
1604 // shambase::get_check_ref(storage.block_cell_sizes)
1605 // .get_buf(id_patch)
1606 // .complete_event_state(resulting_events);
1607 }
1608
1609 void apply_refine_new(
1610 u32 cur_idx,
1611 BlockCoord cur_coords,
1612 std::array<u32, 8> new_blocks,
1613 std::array<BlockCoord, 8> new_block_coords,
1614 RefineCellAccessor acc) const {
1615
1616 auto get_coord_ref = [](u32 i) -> std::array<u32, dim> {
1617 constexpr u32 NsideBlockPow = 1;
1618 constexpr u32 Nside = 1U << NsideBlockPow;
1619
1620 if constexpr (dim == 3) {
1621 const u32 tmp = i >> NsideBlockPow;
1622 return {i % Nside, (tmp) % Nside, (tmp) >> NsideBlockPow};
1623 }
1624 };
1625
1626 auto get_index_block = [](std::array<u32, dim> coord) -> u32 {
1627 constexpr u32 NsideBlockPow = 1;
1628 constexpr u32 Nside = 1U << NsideBlockPow;
1629
1630 if constexpr (dim == 3) {
1631 return coord[0] + Nside * coord[1] + Nside * Nside * coord[2];
1632 }
1633 };
1634
1635 auto get_gid_write = [&](std::array<u32, dim> &glid) -> u32 {
1636 // First, get the block id (it's the block to be refine) in wich the new cell glid
1637 // is located.
1638 std::array<u32, dim> bid
1639 = {glid[0] >> AMRBlock::NsideBlockPow,
1640 glid[1] >> AMRBlock::NsideBlockPow,
1641 glid[2] >> AMRBlock::NsideBlockPow};
1642
1643 // get the new global block id
1644 auto new_glob_id = new_blocks[get_index_block(bid)] * AMRBlock::block_size;
1645
1646 // then added to new_glob_id the local index (between 0 and 7) of the generated
1647 // cells to get. This give the global ids of the new generated cells.
1648 return new_glob_id
1649 + AMRBlock::get_index(
1650 {glid[0] % AMRBlock::Nside,
1651 glid[1] % AMRBlock::Nside,
1652 glid[2] % AMRBlock::Nside});
1653 };
1654
1655 std::array<f64, AMRBlock::block_size> old_rho_block;
1656 std::array<f64_3, AMRBlock::block_size> old_rho_vel_block;
1657 std::array<f64, AMRBlock::block_size> old_rhoE_block;
1658
1659 // save old block
1660 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
1661
1662 auto [lx, ly, lz] = get_coord_ref(loc_id);
1663 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
1664 old_rho_block[loc_id] = acc.rho[old_cell_idx];
1665 old_rho_vel_block[loc_id] = acc.rho_vel[old_cell_idx];
1666 old_rhoE_block[loc_id] = acc.rhoE[old_cell_idx];
1667 }
1668
1669 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
1670
1671 auto [lx, ly, lz] = get_coord_ref(loc_id);
1672 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
1673
1674 // // // cell size in the refined block
1675 // // Tscal delta_cell = cell_sizes[cur_idx];
1676 // // Tscal c_offset = delta_cell * 0.25;
1677 // // std::array<f64_3, AMRBlock::block_size> child_center_offsets;
1678 // // child_center_offsets[0] = {-c_offset, -c_offset, -c_offset}; /*(0,0,0) */
1679 // // child_center_offsets[1] = {c_offset, -c_offset, -c_offset}; /*(1,0,0)*/
1680 // // child_center_offsets[2] = {-c_offset, c_offset, -c_offset}; /* (0,1,0)*/
1681 // // child_center_offsets[3] = {c_offset, c_offset, -c_offset}; /*(1,1,0)*/
1682 // // child_center_offsets[4] = {-c_offset, -c_offset, c_offset}; /*(0,0,1)*/
1683 // // child_center_offsets[5] = {c_offset, -c_offset, c_offset}; /*(1,0,1)*/
1684 // // child_center_offsets[6] = {-c_offset, c_offset, c_offset}; /*(0,1,1)*/
1685 // // child_center_offsets[7] = {c_offset, c_offset, c_offset}; /*(1,1,1)*/
1686
1687 // auto cons_var_slopes = get_3d_grad_cons<Tvec, Minmod>(
1688 // old_cell_idx,
1689 // delta_cell,
1690 // cell_graph_xp,
1691 // cell_graph_xm,
1692 // cell_graph_yp,
1693 // cell_graph_ym,
1694 // cell_graph_zp,
1695 // cell_graph_zm,
1696 // [=](u32 id){
1697 // return acc.rho[id];
1698 // },
1699 // [=](u32 id){
1700 // return acc.rho_vel[id];
1701 // },
1702 // [=](u32 id){
1703 // return acc.rhoE[id];
1704 // });
1705
1706 Tscal rho_block = old_rho_block[loc_id];
1707 Tvec rho_vel_block = old_rho_vel_block[loc_id];
1708 Tscal rhoE_block = old_rhoE_block[loc_id];
1709 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
1710
1711 auto [sx, sy, sz] = get_coord_ref(subdiv_lid);
1712
1713 std::array<u32, 3> glid = {lx * 2 + sx, ly * 2 + sy, lz * 2 + sz};
1714
1715 u32 new_cell_idx = get_gid_write(glid);
1716
1717 // shammath::ConsState<Tvec> cons_var_interp =
1718 // child_center_offsets[subdiv_lid][0] * cons_var_slopes[0] +
1719 // child_center_offsets[subdiv_lid][1] * cons_var_slopes[1] +
1720 // child_center_offsets[subdiv_lid][2] * cons_var_slopes[2];
1721
1722 // // acc.rho[new_cell_idx] = rho_block + cons_var_interp.rho ;
1723 // // acc.rho_vel[new_cell_idx] = rho_vel_block + cons_var_interp.rhovel;
1724 // // acc.rhoE[new_cell_idx] = rhoE_block + cons_var_interp.rhoe;
1725
1726 acc.rho[new_cell_idx] = rho_block;
1727 acc.rho_vel[new_cell_idx] = rho_vel_block;
1728 acc.rhoE[new_cell_idx] = rhoE_block;
1729 }
1730 }
1731 }
1732
1733 void apply_derefine_new(
1734 std::array<u32, 8> old_blocks,
1735 std::array<BlockCoord, 8> old_coords,
1736 u32 new_cell,
1737 BlockCoord new_coord,
1738
1739 RefineCellAccessor acc) const {
1740
1741 std::array<f64, AMRBlock::block_size> rho_block;
1742 std::array<f64_3, AMRBlock::block_size> rho_vel_block;
1743 std::array<f64, AMRBlock::block_size> rhoE_block;
1744
1745 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1746 rho_block[cell_id] = {};
1747 rho_vel_block[cell_id] = {};
1748 rhoE_block[cell_id] = {};
1749 }
1750
1751 // for each siblings block, perform restriction from its 8 children cells
1752 for (u32 pid = 0; pid < 8; pid++) {
1753 auto rho_pid = rho_block[pid];
1754 auto rho_vel_pid = rho_vel_block[pid];
1755 auto rhoe_pid = rhoE_block[pid];
1756
1757 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1758 rho_pid += acc.rho[old_blocks[pid] * AMRBlock::block_size + cell_id];
1759 rho_vel_pid += acc.rho_vel[old_blocks[pid] * AMRBlock::block_size + cell_id];
1760 rhoe_pid += acc.rhoE[old_blocks[pid] * AMRBlock::block_size + cell_id];
1761 }
1762 rho_block[pid] = rho_pid * (1. / 8.);
1763 rho_vel_block[pid] = rho_vel_pid * (1. / 8.);
1764 rhoE_block[pid] = rhoe_pid * (1. / 8.);
1765 }
1766
1767 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1768 u32 newcell_idx = new_cell * AMRBlock::block_size + cell_id;
1769 acc.rho[newcell_idx] = rho_block[cell_id];
1770 acc.rho_vel[newcell_idx] = rho_vel_block[cell_id];
1771 acc.rhoE[newcell_idx] = rhoE_block[cell_id];
1772 }
1773 }
1774 };
1775
1779 class RefineCritPseudoGradientAccessor {
1780 public:
1781 Tscal one_over_Nside = 1. / AMRBlock::Nside;
1782 const TgridVec *block_low_bound;
1783 const TgridVec *block_high_bound;
1784 const Tscal *block_rho;
1785 const f64 *block_pressure;
1786 const f64_3 *block_velocity;
1787
1788 const Tscal *rho_cons;
1789 Tscal error_min;
1790 Tscal error_max;
1791 u32 nblock_per_patch;
1792
1793 AMRGraphLinkiterator cell_graph_xp;
1794 AMRGraphLinkiterator cell_graph_xm;
1795 AMRGraphLinkiterator cell_graph_yp;
1796 AMRGraphLinkiterator cell_graph_ym;
1797 AMRGraphLinkiterator cell_graph_zp;
1798 AMRGraphLinkiterator cell_graph_zm;
1799
1800 RefineCritPseudoGradientAccessor(
1801 sham::EventList &depends_list,
1802 Storage &storage,
1803 u64 id_patch,
1806 Tscal err_min,
1807 Tscal err_max)
1808 : error_min(err_min), error_max(err_max),
1809 cell_graph_xp(
1810 shambase::get_check_ref(storage.cell_graph_edge)
1811 .get_refs_dir(Direction_::xp)
1812 .get(id_patch)
1813 .get()
1814 .get_read_access(depends_list)),
1815 cell_graph_xm(
1816 shambase::get_check_ref(storage.cell_graph_edge)
1817 .get_refs_dir(Direction_::xm)
1818 .get(id_patch)
1819 .get()
1820 .get_read_access(depends_list)),
1821 cell_graph_yp(
1822 shambase::get_check_ref(storage.cell_graph_edge)
1823 .get_refs_dir(Direction_::yp)
1824 .get(id_patch)
1825 .get()
1826 .get_read_access(depends_list)),
1827 cell_graph_ym(
1828 shambase::get_check_ref(storage.cell_graph_edge)
1829 .get_refs_dir(Direction_::ym)
1830 .get(id_patch)
1831 .get()
1832 .get_read_access(depends_list)),
1833 cell_graph_zp(
1834 shambase::get_check_ref(storage.cell_graph_edge)
1835 .get_refs_dir(Direction_::zp)
1836 .get(id_patch)
1837 .get()
1838 .get_read_access(depends_list)),
1839 cell_graph_zm(
1840 shambase::get_check_ref(storage.cell_graph_edge)
1841 .get_refs_dir(Direction_::zm)
1842 .get(id_patch)
1843 .get()
1844 .get_read_access(depends_list))
1845
1846 {
1847 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
1848 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
1849 rho_cons = pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1850 .get_buf()
1851 .get_read_access(depends_list);
1852
1853 nblock_per_patch = pdat.get_obj_cnt();
1854
1855 block_rho = shambase::get_check_ref(storage.rho_primitive)
1856 .get_buf(id_patch)
1857 .get_read_access(depends_list);
1858 block_pressure = shambase::get_check_ref(storage.press)
1859 .get_buf(id_patch)
1860 .get_read_access(depends_list);
1861 block_velocity = shambase::get_check_ref(storage.vel)
1862 .get_buf(id_patch)
1863 .get_read_access(depends_list);
1864 }
1865
1866 void finalize_new(
1867 sham::EventList &resulting_events,
1868 Storage &storage,
1869 u64 id_patch,
1872 Tscal err_min,
1873 Tscal err_max) {
1874
1875 pdat.get_field<i64_3>(0).get_buf().complete_event_state(resulting_events);
1876 pdat.get_field<i64_3>(1).get_buf().complete_event_state(resulting_events);
1877 pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1878 .get_buf()
1879 .complete_event_state(resulting_events);
1880
1881 shambase::get_check_ref(storage.cell_graph_edge)
1882 .get_refs_dir(Direction_::xp)
1883 .get(id_patch)
1884 .get()
1885 .complete_event_state(resulting_events);
1886
1887 shambase::get_check_ref(storage.cell_graph_edge)
1888 .get_refs_dir(Direction_::xm)
1889 .get(id_patch)
1890 .get()
1891 .complete_event_state(resulting_events);
1892
1893 shambase::get_check_ref(storage.cell_graph_edge)
1894 .get_refs_dir(Direction_::yp)
1895 .get(id_patch)
1896 .get()
1897 .complete_event_state(resulting_events);
1898 shambase::get_check_ref(storage.cell_graph_edge)
1899 .get_refs_dir(Direction_::ym)
1900 .get(id_patch)
1901 .get()
1902 .complete_event_state(resulting_events);
1903 shambase::get_check_ref(storage.cell_graph_edge)
1904 .get_refs_dir(Direction_::zp)
1905 .get(id_patch)
1906 .get()
1907 .complete_event_state(resulting_events);
1908 shambase::get_check_ref(storage.cell_graph_edge)
1909 .get_refs_dir(Direction_::zm)
1910 .get(id_patch)
1911 .get()
1912 .complete_event_state(resulting_events);
1913
1914 shambase::get_check_ref(storage.rho_primitive)
1915 .get_buf(id_patch)
1916 .complete_event_state(resulting_events);
1917
1918 shambase::get_check_ref(storage.press)
1919 .get_buf(id_patch)
1920 .complete_event_state(resulting_events);
1921
1922 shambase::get_check_ref(storage.vel)
1923 .get_buf(id_patch)
1924 .complete_event_state(resulting_events);
1925 }
1926
1927 void refine_criterion_new(
1928 u32 block_id,
1929 RefineCritPseudoGradientAccessor acc,
1930 bool &should_refine,
1931 bool &should_derefine) const {
1932 TgridVec low_bound = acc.block_low_bound[block_id];
1933 TgridVec high_bound = acc.block_high_bound[block_id];
1934
1935 // Tscal block_rho_slope = shambase::VectorProperties<Tscal>::get_zero();
1936 // for (u32 i = 0; i < AMRBlock::block_size; i++) {
1937 // block_rho_slope = sham::details::g_sycl_max(
1938 // block_rho_slope,
1939 // modif_second_derivative<Tscal, Tvec>(
1940 // i + block_id * AMRBlock::block_size,
1941 // cell_graph_xp,
1942 // cell_graph_xm,
1943 // cell_graph_yp,
1944 // cell_graph_ym,
1945 // cell_graph_zp,
1946 // cell_graph_zm,
1947 // [=](u32 id) {
1948 // return acc.block_rho[id];
1949 // }));
1950 // }
1951
1952 Tscal block_rho_slope = shambase::VectorProperties<Tscal>::get_zero();
1953 // for (u32 i = 0; i < AMRBlock::block_size; i++) {
1954 // block_rho_slope = sham::details::g_sycl_max(
1955 // block_rho_slope,
1956 // // baryonic_normalized_slope_criterion<Tscal>
1957 // // get_pseudo_grad<Tscal, Tvec>
1958 // (
1959 // i + block_id * AMRBlock::block_size,
1960 // cell_graph_xp,
1961 // cell_graph_xm,
1962 // cell_graph_yp,
1963 // cell_graph_ym,
1964 // cell_graph_zp,
1965 // cell_graph_zm,
1966 // [=](u32 id) {
1967 // // return rho_cons[id];
1968 // return acc.block_rho[id];
1969 // }));
1970 // }
1971
1972 Tscal block_press_grad = shambase::VectorProperties<Tscal>::get_zero();
1973 for (u32 i = 0; i < AMRBlock::block_size; i++) {
1974 block_press_grad = sham::details::g_sycl_max(
1975 block_press_grad,
1976 get_pseudo_grad<Tscal, Tvec>(
1977 i + block_id * AMRBlock::block_size,
1978 cell_graph_xp,
1979 cell_graph_xm,
1980 cell_graph_yp,
1981 cell_graph_ym,
1982 cell_graph_zp,
1983 cell_graph_zm,
1984 [=, this](u32 id) {
1985 return block_pressure[id];
1986 }));
1987 }
1988
1989 Tscal error = sham::details::g_sycl_max(
1990 block_press_grad, sham::details::g_sycl_max(block_rho_slope, 0.0));
1991
1992 should_refine = false;
1993 should_derefine = false;
1994 if (error > error_max) {
1995 should_refine = true;
1996 } else if (error < (error_min * error_max)) {
1997 should_derefine = true;
1998 }
1999
2000 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
2001 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
2002 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
2003 }
2004 };
2005
2009 class RefineCritShearAccessor {
2010 public:
2011 Tscal one_over_Nside = 1. / AMRBlock::Nside;
2012 const TgridVec *block_low_bound;
2013 const TgridVec *block_high_bound;
2014 const Tscal *block_rho;
2015 const f64 *block_pressure;
2016 const f64_3 *block_velocity;
2017
2018 Tscal threshold;
2019 Tscal gamma;
2020 Tscal dxfact;
2021
2022 AMRGraphLinkiterator cell_graph_xp;
2023 AMRGraphLinkiterator cell_graph_xm;
2024 AMRGraphLinkiterator cell_graph_yp;
2025 AMRGraphLinkiterator cell_graph_ym;
2026 AMRGraphLinkiterator cell_graph_zp;
2027 AMRGraphLinkiterator cell_graph_zm;
2028
2029 RefineCritShearAccessor(
2030 sham::EventList &depends_list,
2031 Storage &storage,
2032 u64 id_patch,
2035 Tscal threshold,
2036 Tscal gamma,
2037 Tscal dxfact)
2038 : threshold(threshold), gamma(gamma), dxfact(dxfact),
2039 cell_graph_xp(
2040 shambase::get_check_ref(storage.cell_graph_edge)
2041 .get_refs_dir(Direction_::xp)
2042 .get(id_patch)
2043 .get()
2044 .get_read_access(depends_list)),
2045 cell_graph_xm(
2046 shambase::get_check_ref(storage.cell_graph_edge)
2047 .get_refs_dir(Direction_::xm)
2048 .get(id_patch)
2049 .get()
2050 .get_read_access(depends_list)),
2051 cell_graph_yp(
2052 shambase::get_check_ref(storage.cell_graph_edge)
2053 .get_refs_dir(Direction_::yp)
2054 .get(id_patch)
2055 .get()
2056 .get_read_access(depends_list)),
2057 cell_graph_ym(
2058 shambase::get_check_ref(storage.cell_graph_edge)
2059 .get_refs_dir(Direction_::ym)
2060 .get(id_patch)
2061 .get()
2062 .get_read_access(depends_list)),
2063 cell_graph_zp(
2064 shambase::get_check_ref(storage.cell_graph_edge)
2065 .get_refs_dir(Direction_::zp)
2066 .get(id_patch)
2067 .get()
2068 .get_read_access(depends_list)),
2069 cell_graph_zm(
2070 shambase::get_check_ref(storage.cell_graph_edge)
2071 .get_refs_dir(Direction_::zm)
2072 .get(id_patch)
2073 .get()
2074 .get_read_access(depends_list))
2075
2076 {
2077 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
2078 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
2079
2080 block_rho = shambase::get_check_ref(storage.rho_primitive)
2081 .get_buf(id_patch)
2082 .get_read_access(depends_list);
2083 block_pressure = shambase::get_check_ref(storage.press)
2084 .get_buf(id_patch)
2085 .get_read_access(depends_list);
2086 block_velocity = shambase::get_check_ref(storage.vel)
2087 .get_buf(id_patch)
2088 .get_read_access(depends_list);
2089 }
2090
2091 void finalize_new(
2092 sham::EventList &resulting_events,
2093 Storage &storage,
2094 u64 id_patch,
2097 Tscal threshold,
2098 Tscal gamma,
2099 Tscal dxfact) {
2100
2101 pdat.get_field<i64_3>(0).get_buf().complete_event_state(resulting_events);
2102 pdat.get_field<i64_3>(1).get_buf().complete_event_state(resulting_events);
2103
2104 shambase::get_check_ref(storage.cell_graph_edge)
2105 .get_refs_dir(Direction_::xp)
2106 .get(id_patch)
2107 .get()
2108 .complete_event_state(resulting_events);
2109
2110 shambase::get_check_ref(storage.cell_graph_edge)
2111 .get_refs_dir(Direction_::xm)
2112 .get(id_patch)
2113 .get()
2114 .complete_event_state(resulting_events);
2115
2116 shambase::get_check_ref(storage.cell_graph_edge)
2117 .get_refs_dir(Direction_::yp)
2118 .get(id_patch)
2119 .get()
2120 .complete_event_state(resulting_events);
2121 shambase::get_check_ref(storage.cell_graph_edge)
2122 .get_refs_dir(Direction_::ym)
2123 .get(id_patch)
2124 .get()
2125 .complete_event_state(resulting_events);
2126 shambase::get_check_ref(storage.cell_graph_edge)
2127 .get_refs_dir(Direction_::zp)
2128 .get(id_patch)
2129 .get()
2130 .complete_event_state(resulting_events);
2131 shambase::get_check_ref(storage.cell_graph_edge)
2132 .get_refs_dir(Direction_::zm)
2133 .get(id_patch)
2134 .get()
2135 .complete_event_state(resulting_events);
2136
2137 shambase::get_check_ref(storage.rho_primitive)
2138 .get_buf(id_patch)
2139 .complete_event_state(resulting_events);
2140
2141 shambase::get_check_ref(storage.press)
2142 .get_buf(id_patch)
2143 .complete_event_state(resulting_events);
2144
2145 shambase::get_check_ref(storage.vel)
2146 .get_buf(id_patch)
2147 .complete_event_state(resulting_events);
2148 }
2149
2150 void refine_criterion_new(
2151 u32 block_id,
2152 RefineCritShearAccessor acc,
2153 bool &should_refine,
2154 bool &should_derefine) const {
2155 TgridVec low_bound = acc.block_low_bound[block_id];
2156 TgridVec high_bound = acc.block_high_bound[block_id];
2157
2158 Tvec lower_flt = low_bound.template convert<Tscal>() * dxfact;
2159 Tvec upper_flt = high_bound.template convert<Tscal>() * dxfact;
2160
2161 Tvec block_cell_size = (upper_flt - lower_flt) * one_over_Nside;
2162
2163 Tscal block_normalized_shear = shambase::VectorProperties<Tscal>::get_zero();
2164 for (u32 i = 0; i < AMRBlock::block_size; i++) {
2165 auto cell_id = i + block_id * AMRBlock::block_size;
2166 auto cs = sycl::sqrt(gamma * acc.block_pressure[cell_id] / acc.block_rho[cell_id]);
2167 block_normalized_shear = sham::details::g_sycl_max(
2168 block_normalized_shear,
2169 normalized_shear<Tvec>(
2170 cell_id,
2171 cs,
2172 block_cell_size,
2173
2174 cell_graph_xp,
2175 cell_graph_xm,
2176 cell_graph_yp,
2177 cell_graph_ym,
2178 cell_graph_zp,
2179 cell_graph_zm,
2180 [=](u32 id) {
2181 return acc.block_velocity[id];
2182 }));
2183 }
2184 should_refine = false;
2185 should_derefine = false;
2186 if (block_normalized_shear > threshold * threshold) {
2187 should_refine = true;
2188 } else if (block_normalized_shear < 0.25 * threshold * threshold) {
2189 should_derefine = true;
2190 }
2191
2192 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
2193 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
2194 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
2195 }
2196 };
2197
2201 class RefineCellAccessorAutogravity {
2202 public:
2203 f64 *rho;
2204 f64_3 *rho_vel;
2205 f64 *rhoE;
2206 f64 *phi_old;
2207 f64 *phi_new;
2208
2209 u64 p_id;
2210 // f64* cell_sizes;
2211
2212 // this will be needed for interpolation during refinement
2213 AMRGraphLinkiterator cell_graph_xp;
2214 AMRGraphLinkiterator cell_graph_xm;
2215 AMRGraphLinkiterator cell_graph_yp;
2216 AMRGraphLinkiterator cell_graph_ym;
2217 AMRGraphLinkiterator cell_graph_zp;
2218 AMRGraphLinkiterator cell_graph_zm;
2219
2220 RefineCellAccessorAutogravity(
2221 sham::EventList &depends_list,
2222 Storage &storage,
2223 u64 &id_patch,
2225 : cell_graph_xp(
2226 shambase::get_check_ref(storage.cell_graph_edge)
2227 .get_refs_dir(Direction::xp)
2228 .get(id_patch)
2229 .get()
2230 .get_read_access(depends_list)),
2231 cell_graph_xm(
2232 shambase::get_check_ref(storage.cell_graph_edge)
2233 .get_refs_dir(Direction::xm)
2234 .get(id_patch)
2235 .get()
2236 .get_read_access(depends_list)),
2237 cell_graph_yp(
2238 shambase::get_check_ref(storage.cell_graph_edge)
2239 .get_refs_dir(Direction::yp)
2240 .get(id_patch)
2241 .get()
2242 .get_read_access(depends_list)),
2243 cell_graph_ym(
2244 shambase::get_check_ref(storage.cell_graph_edge)
2245 .get_refs_dir(Direction::ym)
2246 .get(id_patch)
2247 .get()
2248 .get_read_access(depends_list)),
2249 cell_graph_zp(
2250 shambase::get_check_ref(storage.cell_graph_edge)
2251 .get_refs_dir(Direction::zp)
2252 .get(id_patch)
2253 .get()
2254 .get_read_access(depends_list)),
2255 cell_graph_zm(
2256 shambase::get_check_ref(storage.cell_graph_edge)
2257 .get_refs_dir(Direction::zm)
2258 .get(id_patch)
2259 .get()
2260 .get_read_access(depends_list))
2261
2262 {
2263 p_id = id_patch;
2264 rho = pdat.get_field<f64>(2).get_buf().get_write_access(depends_list);
2265 rho_vel = pdat.get_field<f64_3>(3).get_buf().get_write_access(depends_list);
2266 rhoE = pdat.get_field<f64>(4).get_buf().get_write_access(depends_list);
2267 phi_old = pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi_old"))
2268 .get_buf()
2269 .get_write_access(depends_list);
2270 phi_new = pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi"))
2271 .get_buf()
2272 .get_write_access(depends_list);
2273 }
2274
2275 void finalize_new(
2276 sham::EventList &resulting_events,
2277 Storage &storage,
2278 u64 &id_patch,
2280 pdat.get_field<f64>(2).get_buf().complete_event_state(resulting_events);
2281 pdat.get_field<f64_3>(3).get_buf().complete_event_state(resulting_events);
2282 pdat.get_field<f64>(4).get_buf().complete_event_state(resulting_events);
2283 pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi_old"))
2284 .get_buf()
2285 .complete_event_state(resulting_events);
2286 pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi"))
2287 .get_buf()
2288 .complete_event_state(resulting_events);
2289
2290 shambase::get_check_ref(storage.cell_graph_edge)
2291 .get_refs_dir(Direction_::xp)
2292 .get(id_patch)
2293 .get()
2294 .complete_event_state(resulting_events);
2295 shambase::get_check_ref(storage.cell_graph_edge)
2296 .get_refs_dir(Direction_::xm)
2297 .get(id_patch)
2298 .get()
2299 .complete_event_state(resulting_events);
2300 shambase::get_check_ref(storage.cell_graph_edge)
2301 .get_refs_dir(Direction_::yp)
2302 .get(id_patch)
2303 .get()
2304 .complete_event_state(resulting_events);
2305 shambase::get_check_ref(storage.cell_graph_edge)
2306 .get_refs_dir(Direction_::ym)
2307 .get(id_patch)
2308 .get()
2309 .complete_event_state(resulting_events);
2310 shambase::get_check_ref(storage.cell_graph_edge)
2311 .get_refs_dir(Direction_::zp)
2312 .get(id_patch)
2313 .get()
2314 .complete_event_state(resulting_events);
2315 shambase::get_check_ref(storage.cell_graph_edge)
2316 .get_refs_dir(Direction_::zm)
2317 .get(id_patch)
2318 .get()
2319 .complete_event_state(resulting_events);
2320 }
2321
2322 void apply_refine_new(
2323 u32 cur_idx,
2324 BlockCoord cur_coords,
2325 std::array<u32, 8> new_blocks,
2326 std::array<BlockCoord, 8> new_block_coords,
2327 RefineCellAccessorAutogravity acc) const {
2328
2329 auto get_coord_ref = [](u32 i) -> std::array<u32, dim> {
2330 constexpr u32 NsideBlockPow = 1;
2331 constexpr u32 Nside = 1U << NsideBlockPow;
2332
2333 if constexpr (dim == 3) {
2334 const u32 tmp = i >> NsideBlockPow;
2335 return {i % Nside, (tmp) % Nside, (tmp) >> NsideBlockPow};
2336 }
2337 };
2338
2339 auto get_index_block = [](std::array<u32, dim> coord) -> u32 {
2340 constexpr u32 NsideBlockPow = 1;
2341 constexpr u32 Nside = 1U << NsideBlockPow;
2342
2343 if constexpr (dim == 3) {
2344 return coord[0] + Nside * coord[1] + Nside * Nside * coord[2];
2345 }
2346 };
2347
2348 auto get_gid_write = [&](std::array<u32, dim> &glid) -> u32 {
2349 // First, get the block id (it's the block to be refine) in wich the new cell glid
2350 // is located.
2351 std::array<u32, dim> bid
2352 = {glid[0] >> AMRBlock::NsideBlockPow,
2353 glid[1] >> AMRBlock::NsideBlockPow,
2354 glid[2] >> AMRBlock::NsideBlockPow};
2355
2356 // get the new global block id
2357 auto new_glob_id = new_blocks[get_index_block(bid)] * AMRBlock::block_size;
2358
2359 // then added to new_glob_id the local index (between 0 and 7) of the generated
2360 // cells to get. This give the global ids of the new generated cells.
2361 return new_glob_id
2362 + AMRBlock::get_index(
2363 {glid[0] % AMRBlock::Nside,
2364 glid[1] % AMRBlock::Nside,
2365 glid[2] % AMRBlock::Nside});
2366 };
2367
2368 std::array<f64, AMRBlock::block_size> old_rho_block;
2369 std::array<f64_3, AMRBlock::block_size> old_rho_vel_block;
2370 std::array<f64, AMRBlock::block_size> old_rhoE_block;
2371 std::array<f64, AMRBlock::block_size> old_phi_old_block;
2372 std::array<f64, AMRBlock::block_size> old_phi_new_block;
2373
2374 // save old block
2375 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
2376
2377 auto [lx, ly, lz] = get_coord_ref(loc_id);
2378 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
2379 old_rho_block[loc_id] = acc.rho[old_cell_idx];
2380 old_rho_vel_block[loc_id] = acc.rho_vel[old_cell_idx];
2381 old_rhoE_block[loc_id] = acc.rhoE[old_cell_idx];
2382 old_phi_old_block[loc_id] = acc.phi_old[old_cell_idx];
2383 old_phi_new_block[loc_id] = acc.phi_new[old_cell_idx];
2384 }
2385
2386 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
2387
2388 auto [lx, ly, lz] = get_coord_ref(loc_id);
2389 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
2390
2391 // // // cell size in the refined block
2392 // // Tscal delta_cell = cell_sizes[cur_idx];
2393 // // Tscal c_offset = delta_cell * 0.25;
2394 // // std::array<f64_3, AMRBlock::block_size> child_center_offsets;
2395 // // child_center_offsets[0] = {-c_offset, -c_offset, -c_offset}; /*(0,0,0) */
2396 // // child_center_offsets[1] = {c_offset, -c_offset, -c_offset}; /*(1,0,0)*/
2397 // // child_center_offsets[2] = {-c_offset, c_offset, -c_offset}; /* (0,1,0)*/
2398 // // child_center_offsets[3] = {c_offset, c_offset, -c_offset}; /*(1,1,0)*/
2399 // // child_center_offsets[4] = {-c_offset, -c_offset, c_offset}; /*(0,0,1)*/
2400 // // child_center_offsets[5] = {c_offset, -c_offset, c_offset}; /*(1,0,1)*/
2401 // // child_center_offsets[6] = {-c_offset, c_offset, c_offset}; /*(0,1,1)*/
2402 // // child_center_offsets[7] = {c_offset, c_offset, c_offset}; /*(1,1,1)*/
2403
2404 // auto cons_var_slopes = get_3d_grad_cons<Tvec, Minmod>(
2405 // old_cell_idx,
2406 // delta_cell,
2407 // cell_graph_xp,
2408 // cell_graph_xm,
2409 // cell_graph_yp,
2410 // cell_graph_ym,
2411 // cell_graph_zp,
2412 // cell_graph_zm,
2413 // [=](u32 id){
2414 // return acc.rho[id];
2415 // },
2416 // [=](u32 id){
2417 // return acc.rho_vel[id];
2418 // },
2419 // [=](u32 id){
2420 // return acc.rhoE[id];
2421 // });
2422
2423 Tscal rho_block = old_rho_block[loc_id];
2424 Tvec rho_vel_block = old_rho_vel_block[loc_id];
2425 Tscal rhoE_block = old_rhoE_block[loc_id];
2426 Tscal phi_old_block = old_phi_old_block[loc_id];
2427 Tscal phi_new_block = old_phi_new_block[loc_id];
2428
2429 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
2430
2431 auto [sx, sy, sz] = get_coord_ref(subdiv_lid);
2432
2433 std::array<u32, 3> glid = {lx * 2 + sx, ly * 2 + sy, lz * 2 + sz};
2434
2435 u32 new_cell_idx = get_gid_write(glid);
2436
2437 // shammath::ConsState<Tvec> cons_var_interp =
2438 // child_center_offsets[subdiv_lid][0] * cons_var_slopes[0] +
2439 // child_center_offsets[subdiv_lid][1] * cons_var_slopes[1] +
2440 // child_center_offsets[subdiv_lid][2] * cons_var_slopes[2];
2441
2442 // // acc.rho[new_cell_idx] = rho_block + cons_var_interp.rho ;
2443 // // acc.rho_vel[new_cell_idx] = rho_vel_block + cons_var_interp.rhovel;
2444 // // acc.rhoE[new_cell_idx] = rhoE_block + cons_var_interp.rhoe;
2445
2446 acc.rho[new_cell_idx] = rho_block;
2447 acc.rho_vel[new_cell_idx] = rho_vel_block;
2448 acc.rhoE[new_cell_idx] = rhoE_block;
2449 acc.phi_old[new_cell_idx] = phi_old_block;
2450 acc.phi_new[new_cell_idx] = phi_new_block;
2451 }
2452 }
2453 }
2454
2455 void apply_derefine_new(
2456 std::array<u32, 8> old_blocks,
2457 std::array<BlockCoord, 8> old_coords,
2458 u32 new_cell,
2459 BlockCoord new_coord,
2460
2461 RefineCellAccessorAutogravity acc) const {
2462
2463 std::array<f64, AMRBlock::block_size> rho_block;
2464 std::array<f64_3, AMRBlock::block_size> rho_vel_block;
2465 std::array<f64, AMRBlock::block_size> rhoE_block;
2466 std::array<f64, AMRBlock::block_size> phi_old_block;
2467 std::array<f64, AMRBlock::block_size> phi_new_block;
2468
2469 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2470 rho_block[cell_id] = {};
2471 rho_vel_block[cell_id] = {};
2472 rhoE_block[cell_id] = {};
2473 phi_old_block[cell_id] = {};
2474 phi_new_block[cell_id] = {};
2475 }
2476
2477 // for each siblings block, perform restriction from its 8 children cells
2478 for (u32 pid = 0; pid < 8; pid++) {
2479 auto rho_pid = rho_block[pid];
2480 auto rho_vel_pid = rho_vel_block[pid];
2481 auto rhoe_pid = rhoE_block[pid];
2482 auto phi_old_pid = phi_old_block[pid];
2483 auto phi_new_pid = phi_new_block[pid];
2484
2485 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2486 rho_pid += acc.rho[old_blocks[pid] * AMRBlock::block_size + cell_id];
2487 rho_vel_pid += acc.rho_vel[old_blocks[pid] * AMRBlock::block_size + cell_id];
2488 rhoe_pid += acc.rhoE[old_blocks[pid] * AMRBlock::block_size + cell_id];
2489 phi_old_pid += acc.phi_old[old_blocks[pid] * AMRBlock::block_size + cell_id];
2490 phi_new_pid += acc.phi_new[old_blocks[pid] * AMRBlock::block_size + cell_id];
2491 }
2492 rho_block[pid] = rho_pid * (1. / 8.);
2493 rho_vel_block[pid] = rho_vel_pid * (1. / 8.);
2494 rhoE_block[pid] = rhoe_pid * (1. / 8.);
2495 // phi_old_block[pid] = phi_old_pid * (1. / 8.);
2496 // phi_new_block[pid] = phi_new_pid * (1. / 8.);
2497 }
2498
2499 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2500 u32 newcell_idx = new_cell * AMRBlock::block_size + cell_id;
2501 acc.rho[newcell_idx] = rho_block[cell_id];
2502 acc.rho_vel[newcell_idx] = rho_vel_block[cell_id];
2503 acc.rhoE[newcell_idx] = rhoE_block[cell_id];
2504
2505 // acc.phi_old[newcell_idx] = phi_old_block[cell_id];
2506 // acc.phi_new[newcell_idx] = phi_new_block[cell_id];
2507 }
2508 }
2509 };
2510
2511 using AMRmode_None = typename AMRMode<Tvec, TgridVec>::None;
2512 using AMRmode_DensityBased = typename AMRMode<Tvec, TgridVec>::DensityBased;
2513 using AMRmode_PseudoGradientBased = typename AMRMode<Tvec, TgridVec>::PseudoGradientBased;
2514 using AMRmode_JeansLengthBased = typename AMRMode<Tvec, TgridVec>::JeansLengthBased;
2515 using AMRmode_ShearBased = typename AMRMode<Tvec, TgridVec>::ShearBased;
2516
2517 bool has_cell_order_changed = false;
2518
2519 // get refine and derefine list
2522
2523 if (AMRmode_None *cfg = std::get_if<AMRmode_None>(&solver_config.amr_mode.config)) {
2524 // no refinment here turn around there is nothing to see
2525 } else {
2526 if (AMRmode_DensityBased *cfg
2527 = std::get_if<AMRmode_DensityBased>(&solver_config.amr_mode.config)) {
2528
2529 Tscal dxfact(solver_config.grid_coord_to_pos_fact);
2530 gen_refine_block_changes_new<RefineCritBlock>(
2531 refine_list, derefine_list, dxfact, cfg->crit_mass);
2532 }
2533
2534 else if (
2535 AMRmode_PseudoGradientBased *cfg
2536 = std::get_if<AMRmode_PseudoGradientBased>(&solver_config.amr_mode.config)) {
2537
2538 gen_refine_block_changes_new<RefineCritPseudoGradientAccessor>(
2539 refine_list, derefine_list, cfg->error_min, cfg->error_max);
2540 }
2541
2542 else if (
2543 AMRmode_ShearBased *cfg
2544 = std::get_if<AMRmode_ShearBased>(&solver_config.amr_mode.config)) {
2545 Tscal dxfact(solver_config.grid_coord_to_pos_fact);
2546 Tscal gamma(solver_config.eos_gamma);
2547
2548 gen_refine_block_changes_new<RefineCritShearAccessor>(
2549 refine_list, derefine_list, cfg->threshold, gamma, dxfact);
2550 }
2551
2553 enforce_two_to_one_refinement_new(std::move(refine_list));
2555 enforce_two_to_one_derefinement_new(std::move(derefine_list), std::move(refine_list));
2557 // Note that this only add new blocks at the end of the patchdata
2558 bool change_refine = internal_refine_grid_new<RefineCellAccessor>(std::move(refine_list));
2559
2561 // Note that this will perform the merge then remove the old blocks
2562 // This is ok to call straight after the refine without edditing the index list in
2563 // derefine_list since no permutations were applied in internal_refine_grid_new and no cells
2564 // can be both refined and derefined in the same pass
2565 bool change_derefine
2566 = internal_derefine_grid_new<RefineCellAccessor>(std::move(derefine_list));
2567
2568 has_cell_order_changed = has_cell_order_changed || (change_refine || change_derefine);
2569
2570 if (has_cell_order_changed) {
2571 // Ensure that the blocks are sorted before refinement
2572 AMRSortBlocks block_sorter(context, solver_config, storage);
2573 block_sorter.reorder_amr_blocks();
2574 }
2575 }
2576}
2577
double f64
Alias for double.
std::uint32_t u32
32 bit unsigned integer
std::uint64_t u64
64 bit unsigned integer
A buffer allocated in USM (Unified Shared Memory).
void complete_event_state(sycl::event e) const
Complete the event state of the buffer.
T * get_write_access(sham::EventList &depends_list, SourceLocation src_loc=SourceLocation{})
Get a read-write pointer to the buffer's data.
void copy_range(size_t begin, size_t end, sham::DeviceBuffer< T, dest_target > &dest) const
Copy a range of elements from the buffer to another buffer.
size_t get_size() const
Gets the number of elements in the buffer.
const T * get_read_access(sham::EventList &depends_list, SourceLocation src_loc=SourceLocation{}) const
Get a read-only pointer to the buffer's data.
A SYCL queue associated with a device and a context.
sycl::queue q
The SYCL queue associated with this context.
sycl::event submit(Fct &&fct)
Submits a kernel to the SYCL queue.
Class to manage a list of SYCL events.
Definition EventList.hpp:31
void add_event(sycl::event e)
Add an event to the list of events.
Definition EventList.hpp:87
Represents a collection of objects distributed across patches identified by a u64 id.
PatchDataLayer container class, the layout is described in patchdata_layout.
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...
main include file for the shamalgs algorithms
Slope mode enum + json serialization/deserialization.
alias namespace to simplify the use of log functions
Definition logs.hpp:319
sycl::buffer< typename std::invoke_result_t< Fct, u32 > > gen_buffer_device(sycl::queue &q, u32 len, Fct &&func)
generate a buffer from a lambda expression based on the indexes
Definition algorithm.hpp:65
std::tuple< std::optional< sycl::buffer< u32 > >, u32 > stream_compact(sycl::queue &q, sycl::buffer< u32 > &buf_flags, u32 len)
Stream compaction algorithm.
Definition numeric.cpp:84
constexpr T pow_constexpr(T a) noexcept
Calculates the power of a number at compile time.
Definition integer.hpp:174
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
std::vector< std::string_view > args
Executable argument list (mapped from argv).
Definition cmdopt.cpp:63
void raw_ln(Types... var2)
Prints a log message with multiple arguments followed by a newline.
Definition logs.hpp:90
void info_ln(std::string module_name, Types... var2)
Prints a log message with multiple arguments followed by a newline.
Definition logs.hpp:133
Patch object that contain generic patch information.
Definition Patch.hpp:33
u64 id_patch
unique key that identify the patch
Definition Patch.hpp:86