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
973 if constexpr (dim == 3) {
974 const u32 tmp = i >> NsideBlockPow;
975 // This line is why derefinement never happens
976 // return {i % Nside, (tmp) % Nside, (tmp ) >> NsideBlockPow};
977 return {(tmp) >> NsideBlockPow, (tmp) % Nside, i % Nside};
978 }
979 };
980
981 auto get_split
982 = [=](BlockCoord target_block) -> std::array<BlockCoord, split_count> {
983 std::array<BlockCoord, split_count> ret;
984 auto bmin = target_block.bmin;
985 auto bmax = target_block.bmax;
986 auto split = bmin + (bmax - bmin) / 2;
987 std::array<TgridVec, 3> szs = {bmin, split, bmax};
988 for (u32 i = 0; i < split_count; i++) {
989 auto [lx, ly, lz] = get_coord(i);
990
991 ret[i].bmin = TgridVec{szs[lx].x(), szs[ly].y(), szs[lz].z()};
992 ret[i].bmax
993 = TgridVec{szs[lx + 1].x(), szs[ly + 1].y(), szs[lz + 1].z()};
994 }
995
996 return ret;
997 };
998
999 for (u32 b_lid = 0; b_lid < split_count; b_lid++) {
1000 blocks[b_lid] = BlockCoord{acc_min[id + b_lid], acc_max[id + b_lid]};
1001 all_want_to_merge = all_want_to_merge && acc_merge_flag[id + b_lid];
1002 all_same_level
1003 = all_same_level && (acc_amr_levels[id] == acc_amr_levels[id + b_lid]);
1004 }
1005
1006 BlockCoord merged = BlockCoord::get_merge(blocks);
1007 std::array<BlockCoord, split_count> splitted = get_split(merged);
1008 for (u32 lid = 0; lid < split_count; lid++) {
1009 do_merge = do_merge && sham::equals(blocks[lid].bmin, splitted[lid].bmin)
1010 && sham::equals(blocks[lid].bmax, splitted[lid].bmax);
1011 }
1012
1013 do_merge = do_merge && all_want_to_merge && all_same_level;
1014 if (acc_refine_flag[id] && do_merge) {
1015 do_merge = false;
1016 }
1017
1018 } else {
1019 do_merge = false;
1020 }
1021 acc_merge_flag[id] = do_merge;
1022 });
1023 });
1024 buf_cell_min.complete_event_state(e);
1025 buf_cell_max.complete_event_state(e);
1026 buf_amr_block_levels.complete_event_state(e);
1027 patch_derefine_flag.complete_event_state(e);
1028 patch_refine_flag.complete_event_state(e);
1029
1032 auto buf_derefine_1
1033 = shamalgs::numeric::stream_compact(dev_sched, patch_derefine_flag, obj_cnt);
1035 " Count block's flag for derefinement [After geometry validity check and before 2:1 "
1036 "check] "
1037 "\t : ",
1038 buf_derefine_1.get_size(),
1039 "\n");
1042
1043 // ////////////////////////////////////////////////////////////////////////////////////
1044 // // // enforce 2:1 at parent level
1045 // ///////////////////////////////////////////////////////////////////////////////////
1046
1047 std::shared_ptr<sham::DeviceScheduler> dev_sched
1048 = shamsys::instance::get_compute_scheduler_ptr();
1049
1050 sham::DeviceBuffer<u32> changed_buf(1, dev_sched);
1051
1052 // copy old deref buffer to avoid race condition
1053 sham::DeviceBuffer<u32> patch_derefine_flag_old(obj_cnt, dev_sched);
1054 patch_derefine_flag.copy_range(0, obj_cnt, patch_derefine_flag_old);
1055
1056 //
1057 sham::DeviceBuffer<u32> patch_derefine_flag_new(obj_cnt, dev_sched);
1058
1059 for (int it = 0; it < 100; it++) {
1060 changed_buf.set_val_at_idx(0, 0);
1061
1062 sham::EventList depend_list;
1063
1064 AMRGraphLinkiterator block_graph_xp
1065 = block_graph_neighs_xp.get_read_access(depend_list);
1066
1067 AMRGraphLinkiterator block_graph_xm
1068 = block_graph_neighs_xm.get_read_access(depend_list);
1069 AMRGraphLinkiterator block_graph_yp
1070 = block_graph_neighs_yp.get_read_access(depend_list);
1071 AMRGraphLinkiterator block_graph_ym
1072 = block_graph_neighs_ym.get_read_access(depend_list);
1073 AMRGraphLinkiterator block_graph_zp
1074 = block_graph_neighs_zp.get_read_access(depend_list);
1075 AMRGraphLinkiterator block_graph_zm
1076 = block_graph_neighs_zm.get_read_access(depend_list);
1077
1078 auto acc_amr_levels = buf_amr_block_levels.get_read_access(depend_list);
1079 auto acc_ref_flag = patch_refine_flag.get_read_access(depend_list);
1080 auto acc_changed = changed_buf.get_write_access(depend_list);
1081
1082 auto acc_deref_old = patch_derefine_flag_old.get_read_access(depend_list);
1083 auto acc_deref_new = patch_derefine_flag_new.get_write_access(depend_list);
1084
1085 auto e_2to1 = q.submit(depend_list, [&](sycl::handler &cgh) {
1086 cgh.parallel_for(sycl::range<1>(obj_cnt), [=](sycl::item<1> gid) {
1087 auto lid = gid.get_linear_id();
1088
1089 auto old_flag = acc_deref_old[lid];
1090 auto new_flag = old_flag;
1091
1092 auto check_2To1_der = [&](u32 nid) {
1093 if (nid < obj_cnt)
1094
1095 {
1096
1097 auto neigh_future = acc_amr_levels[nid] + (acc_ref_flag[nid] ? 1 : 0)
1098 - (acc_deref_old[nid] ? 1 : 0);
1099
1100 auto my_future = acc_amr_levels[lid] - 1;
1101
1102 if (neigh_future > my_future + 1) {
1103 new_flag = 0;
1104 }
1105 }
1106 };
1107
1108 if (old_flag) {
1109
1110 for (u32 i = 0; i < AMRBlock::block_size; i++) {
1111 block_graph_xp.for_each_object_link((lid + i), check_2To1_der);
1112 block_graph_xm.for_each_object_link((lid + i), check_2To1_der);
1113 block_graph_yp.for_each_object_link((lid + i), check_2To1_der);
1114 block_graph_ym.for_each_object_link((lid + i), check_2To1_der);
1115 block_graph_zp.for_each_object_link((lid + i), check_2To1_der);
1116 block_graph_zm.for_each_object_link((lid + i), check_2To1_der);
1117 }
1118
1119 if (old_flag != new_flag) {
1120 sycl::atomic_ref<
1121 u32,
1122 sycl::memory_order::relaxed,
1123 sycl::memory_scope::system>
1124 atomic_changed(acc_changed[0]);
1125 atomic_changed.exchange(1);
1126 }
1127 }
1128
1129 acc_deref_new[lid] = new_flag;
1130 });
1131 });
1132 block_graph_neighs_xp.complete_event_state(e_2to1);
1133 block_graph_neighs_xm.complete_event_state(e_2to1);
1134 block_graph_neighs_yp.complete_event_state(e_2to1);
1135 block_graph_neighs_ym.complete_event_state(e_2to1);
1136 block_graph_neighs_zp.complete_event_state(e_2to1);
1137 block_graph_neighs_zm.complete_event_state(e_2to1);
1138 buf_amr_block_levels.complete_event_state(e_2to1);
1139 patch_refine_flag.complete_event_state(e_2to1);
1140 changed_buf.complete_event_state(e_2to1);
1141
1142 patch_derefine_flag_old.complete_event_state(e_2to1);
1143 patch_derefine_flag_new.complete_event_state(e_2to1);
1144 e_2to1.wait();
1145
1146 std::swap(patch_derefine_flag_old, patch_derefine_flag_new);
1147
1148 if (changed_buf.get_val_at_idx(0) == 0) {
1149 logger:
1151 "Derefinement 2:1 balance converge in \t ", it + 1, "\t sweeps \n\n");
1152 break;
1153 }
1154 }
1155
1156 // copy back to ..
1157 patch_derefine_flag_old.copy_range(0, obj_cnt, patch_derefine_flag);
1158
1159 // ////////////////////////////////////////////////////////////////////////////////
1160 // // derefinement
1161 // ////////////////////////////////////////////////////////////////////////////////
1162 // perform stream compactions on the derefinement flags
1163 auto buf_derefine
1164 = shamalgs::numeric::stream_compact(dev_sched, patch_derefine_flag, obj_cnt);
1165
1167 " Count block's flag for derefinement [After geometry validity check and after 2:1 "
1168 "check] \t : ",
1169 buf_derefine.get_size(),
1170 "\n");
1171
1172 shamlog_debug_ln(
1173 "AMRGrid", "patch ", id_patch, buf_derefine.get_size(), "marked for derefinement ");
1174 });
1175}
1176
1177template<class Tvec, class TgridVec>
1178template<class UserAcc>
1182 const AMRInterpMode amr_refine_interp_mode) {
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, amr_refine_interp_mode);
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{e};
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, amr_refine_interp_mode);
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 const AMRInterpMode amr_refine_interp_mode) {
1279
1280 using namespace shamrock::patch;
1281
1282 bool cell_were_removed = false;
1283
1284 sham::DeviceQueue &q = shamsys::instance::get_compute_scheduler().get_queue();
1285 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
1286
1287 scheduler().for_each_patch_data([&](u64 id_patch, Patch cur_p, PatchDataLayer &pdat) {
1288 u32 old_obj_cnt = pdat.get_obj_cnt();
1289 u32 old_obj_cnt_before_refinement = dd_derefine_flags.get(id_patch).get_size();
1290 auto stream_compact_results = shamalgs::numeric::stream_compact(
1291 dev_sched, dd_derefine_flags.get(id_patch), old_obj_cnt_before_refinement);
1292 if (stream_compact_results.get_size() > 0) {
1293 // init flag table
1294 sycl::buffer<u32> keep_block_flag
1295 = shamalgs::algorithm::gen_buffer_device(q.q, old_obj_cnt, [](u32 i) -> u32 {
1296 return 1;
1297 });
1298
1299 sham::DeviceBuffer<TgridVec> &buf_cell_min = pdat.get_field_buf_ref<TgridVec>(0);
1300 sham::DeviceBuffer<TgridVec> &buf_cell_max = pdat.get_field_buf_ref<TgridVec>(1);
1301 sham::EventList depends_list;
1302 auto block_bound_low = buf_cell_min.get_write_access(depends_list);
1303 auto block_bound_high = buf_cell_max.get_write_access(depends_list);
1304 UserAcc uacc(depends_list, storage, id_patch, pdat, amr_refine_interp_mode);
1305 auto index_to_deref = stream_compact_results.get_read_access(depends_list);
1306
1307 // edit block content + make flag of blocks to keep
1308 auto e = q.submit(depends_list, [&](sycl::handler &cgh) {
1309 sycl::accessor flag_keep{keep_block_flag, cgh, sycl::read_write};
1310 cgh.parallel_for(
1311 sycl::range<1>(stream_compact_results.get_size()), [=](sycl::item<1> gid) {
1312 u32 tid = gid.get_linear_id();
1313
1314 u32 idx_to_derefine = index_to_deref[gid];
1315
1316 // compute old block indexes
1317 std::array<u32, split_count> old_indexes;
1318#pragma unroll
1319 for (u32 pid = 0; pid < split_count; pid++) {
1320 old_indexes[pid] = idx_to_derefine + pid;
1321 }
1322
1323 // load block coords
1324 std::array<BlockCoord, split_count> block_coords;
1325#pragma unroll
1326 for (u32 pid = 0; pid < split_count; pid++) {
1327 block_coords[pid] = BlockCoord{
1328 block_bound_low[old_indexes[pid]],
1329 block_bound_high[old_indexes[pid]]};
1330 }
1331
1332 // make new block coord
1333 BlockCoord merged_block_coord = BlockCoord::get_merge(block_coords);
1334
1335 // write new coord
1336 block_bound_low[idx_to_derefine] = merged_block_coord.bmin;
1337 block_bound_high[idx_to_derefine] = merged_block_coord.bmax;
1338
1339// flag the old blocks for removal
1340#pragma unroll
1341 for (u32 pid = 1; pid < split_count; pid++) {
1342 flag_keep[idx_to_derefine + pid] = 0;
1343 }
1344
1345 // user lambda to fill the fields
1346
1347 uacc.apply_derefine_new(
1348 old_indexes, block_coords, idx_to_derefine, merged_block_coord, uacc);
1349 });
1350 });
1351
1352 sham::EventList resulting_events;
1353
1354 buf_cell_min.complete_event_state(resulting_events);
1355 buf_cell_max.complete_event_state(resulting_events);
1356 uacc.finalize_new(resulting_events, storage, id_patch, pdat, amr_refine_interp_mode);
1357
1358 stream_compact_results.complete_event_state(resulting_events);
1359
1360 // stream compact the flags (get new block ids map after merged)
1361 auto [opt_buf, len]
1362 = shamalgs::numeric::stream_compact(q.q, keep_block_flag, old_obj_cnt);
1363
1365 "AMR Grid",
1366 "patch",
1367 id_patch,
1368 "derefine block count = ",
1369 old_obj_cnt - len,
1370 "new block count = ",
1371 len);
1372
1373 if (!opt_buf) {
1374 throw std::runtime_error("opt buf must contain something at this point");
1375 }
1376
1377 // remap pdat according to stream compact (for each field in patchdataleyer resize
1378 // according to new block ids map)
1379 pdat.index_remap_resize(*opt_buf, len);
1380
1381 cell_were_removed = cell_were_removed || stream_compact_results.get_size() > 0;
1382 }
1383 });
1384
1385 return cell_were_removed;
1386}
1387
1388template<class Tvec, class TgridVec>
1389template<class UserAccCrit, class UserAccSplit, class UserAccMerge>
1391 internal_update_refinement_new(const AMRInterpMode amr_refine_interp_mode) {
1392
1393 // Ensure that the blocks are sorted before refinement
1394 AMRSortBlocks block_sorter(context, solver_config, storage);
1395 block_sorter.reorder_amr_blocks();
1396
1397 // get refine and derefine list
1400
1401 gen_refine_block_changes_new<UserAccCrit>(dd_refine_list, dd_derefine_list);
1402
1404 // Note that this only add new blocks at the end of the patchdata
1405 internal_refine_grid_new<UserAccSplit>(std::move(dd_refine_list), amr_refine_interp_mode);
1406
1408 // Note that this will perform the merge then remove the old blocks
1409 // This is ok to call straight after the refine without edditing the index list in derefine_list
1410 // since no permutations were applied in internal_refine_grid_new and no cells can be both
1411 // refined and derefined in the same pass
1412 internal_derefine_grid_new<UserAccMerge>(std::move(dd_derefine_list), amr_refine_interp_mode);
1413}
1414
1415template<class Tvec, class TgridVec>
1418
1419 class RefineCritBlock {
1420 public:
1421 const TgridVec *block_low_bound;
1422 const TgridVec *block_high_bound;
1423 const Tscal *block_density_field;
1424
1425 Tscal one_over_Nside = 1. / AMRBlock::Nside;
1426
1427 Tscal dxfact;
1428 Tscal wanted_mass;
1429
1430 RefineCritBlock(
1431 sham::EventList &depends_list,
1432 Storage &storage,
1433 u64 id_patch,
1436 Tscal dxfact,
1437 Tscal wanted_mass)
1438 : dxfact(dxfact), wanted_mass(wanted_mass) {
1439
1440 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
1441 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
1442 block_density_field = pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1443 .get_buf()
1444 .get_read_access(depends_list);
1445 }
1446
1447 void finalize_new(
1448 sham::EventList &resulting_events,
1449 Storage &storage,
1450 u64 id_patch,
1453 Tscal dxfact,
1454 Tscal wanted_mass) {
1455
1456 pdat.get_field<TgridVec>(0).get_buf().complete_event_state(resulting_events);
1457 pdat.get_field<TgridVec>(1).get_buf().complete_event_state(resulting_events);
1458 pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1459 .get_buf()
1460 .complete_event_state(resulting_events);
1461 }
1462
1463 void refine_criterion_new(
1464 u32 block_id, RefineCritBlock acc, bool &should_refine, bool &should_derefine) const {
1465
1466 TgridVec low_bound = acc.block_low_bound[block_id];
1467 TgridVec high_bound = acc.block_high_bound[block_id];
1468
1469 Tvec lower_flt = low_bound.template convert<Tscal>() * dxfact;
1470 Tvec upper_flt = high_bound.template convert<Tscal>() * dxfact;
1471
1472 Tvec block_cell_size = (upper_flt - lower_flt) * one_over_Nside;
1473
1474 Tscal sum_mass = 0;
1475 for (u32 i = 0; i < AMRBlock::block_size; i++) {
1476 sum_mass += acc.block_density_field[i + block_id * AMRBlock::block_size];
1477 }
1478 sum_mass *= block_cell_size.x() * block_cell_size.y() * block_cell_size.z();
1479
1480 if (sum_mass > wanted_mass * 8) {
1481 should_refine = true;
1482 should_derefine = false;
1483 } else if (sum_mass < wanted_mass) {
1484 should_refine = false;
1485 should_derefine = true;
1486 } else {
1487 should_refine = false;
1488 should_derefine = false;
1489 }
1490
1491 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
1492 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
1493 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
1494 }
1495 };
1496
1497 class RefineCellAccessor {
1498 public:
1499 f64 *rho;
1500 f64_3 *rho_vel;
1501 f64 *rhoE;
1502 u64 p_id;
1503 f64 *cell_sizes;
1504
1505 const f64 *rho_old_snap;
1506 const f64_3 *rho_vel_old_snap;
1507 const f64 *rhoE_old_snap;
1508
1509 AMRInterpMode amr_ref_interp_mode;
1510
1511 // this will be needed for interpolation during refinement
1512 AMRGraphLinkiterator cell_graph_xp;
1513 AMRGraphLinkiterator cell_graph_xm;
1514 AMRGraphLinkiterator cell_graph_yp;
1515 AMRGraphLinkiterator cell_graph_ym;
1516 AMRGraphLinkiterator cell_graph_zp;
1517 AMRGraphLinkiterator cell_graph_zm;
1518
1519 RefineCellAccessor(
1520 sham::EventList &depends_list,
1521 Storage &storage,
1522 u64 &id_patch,
1524 AMRInterpMode _amr_ref_interp_mode)
1525 : cell_graph_xp(
1526 shambase::get_check_ref(storage.cell_graph_edge)
1527 .get_refs_dir(Direction::xp)
1528 .get(id_patch)
1529 .get()
1530 .get_read_access(depends_list)),
1531 cell_graph_xm(
1532 shambase::get_check_ref(storage.cell_graph_edge)
1533 .get_refs_dir(Direction::xm)
1534 .get(id_patch)
1535 .get()
1536 .get_read_access(depends_list)),
1537 cell_graph_yp(
1538 shambase::get_check_ref(storage.cell_graph_edge)
1539 .get_refs_dir(Direction::yp)
1540 .get(id_patch)
1541 .get()
1542 .get_read_access(depends_list)),
1543 cell_graph_ym(
1544 shambase::get_check_ref(storage.cell_graph_edge)
1545 .get_refs_dir(Direction::ym)
1546 .get(id_patch)
1547 .get()
1548 .get_read_access(depends_list)),
1549 cell_graph_zp(
1550 shambase::get_check_ref(storage.cell_graph_edge)
1551 .get_refs_dir(Direction::zp)
1552 .get(id_patch)
1553 .get()
1554 .get_read_access(depends_list)),
1555 cell_graph_zm(
1556 shambase::get_check_ref(storage.cell_graph_edge)
1557 .get_refs_dir(Direction::zm)
1558 .get(id_patch)
1559 .get()
1560 .get_read_access(depends_list)),
1561 amr_ref_interp_mode(_amr_ref_interp_mode)
1562
1563 {
1564 p_id = id_patch;
1565
1566 // old conservatives variables
1567 rho_old_snap = shambase::get_check_ref(storage.rho_snap)
1568 .get(id_patch)
1569 .get_buf()
1570 .get_read_access(depends_list);
1571 rhoE_old_snap = shambase::get_check_ref(storage.rhoe_snap)
1572 .get(id_patch)
1573 .get_buf()
1574 .get_read_access(depends_list);
1575 rho_vel_old_snap = shambase::get_check_ref(storage.rho_vel_snap)
1576 .get(id_patch)
1577 .get_buf()
1578 .get_read_access(depends_list);
1579
1580 // new conservative variables
1581 rho = pdat.get_field<f64>(2).get_buf().get_write_access(depends_list);
1582 rho_vel = pdat.get_field<f64_3>(3).get_buf().get_write_access(depends_list);
1583 rhoE = pdat.get_field<f64>(4).get_buf().get_write_access(depends_list);
1584 cell_sizes = shambase::get_check_ref(storage.block_cell_sizes)
1585 .get_buf(id_patch)
1586 .get_write_access(depends_list);
1587 }
1588
1589 void finalize_new(
1590 sham::EventList &resulting_events,
1591 Storage &storage,
1592 u64 &id_patch,
1594 AMRInterpMode amr_ref_interp_mode) {
1595 pdat.get_field<f64>(2).get_buf().complete_event_state(resulting_events);
1596 pdat.get_field<f64_3>(3).get_buf().complete_event_state(resulting_events);
1597 pdat.get_field<f64>(4).get_buf().complete_event_state(resulting_events);
1598
1599 shambase::get_check_ref(storage.cell_graph_edge)
1600 .get_refs_dir(Direction_::xp)
1601 .get(id_patch)
1602 .get()
1603 .complete_event_state(resulting_events);
1604 shambase::get_check_ref(storage.cell_graph_edge)
1605 .get_refs_dir(Direction_::xm)
1606 .get(id_patch)
1607 .get()
1608 .complete_event_state(resulting_events);
1609 shambase::get_check_ref(storage.cell_graph_edge)
1610 .get_refs_dir(Direction_::yp)
1611 .get(id_patch)
1612 .get()
1613 .complete_event_state(resulting_events);
1614 shambase::get_check_ref(storage.cell_graph_edge)
1615 .get_refs_dir(Direction_::ym)
1616 .get(id_patch)
1617 .get()
1618 .complete_event_state(resulting_events);
1619 shambase::get_check_ref(storage.cell_graph_edge)
1620 .get_refs_dir(Direction_::zp)
1621 .get(id_patch)
1622 .get()
1623 .complete_event_state(resulting_events);
1624 shambase::get_check_ref(storage.cell_graph_edge)
1625 .get_refs_dir(Direction_::zm)
1626 .get(id_patch)
1627 .get()
1628 .complete_event_state(resulting_events);
1629
1630 shambase::get_check_ref(storage.block_cell_sizes)
1631 .get_buf(id_patch)
1632 .complete_event_state(resulting_events);
1633
1634 // old conservative variables
1635 shambase::get_check_ref(storage.rho_snap)
1636 .get(id_patch)
1637 .get_buf()
1638 .complete_event_state(resulting_events);
1639 shambase::get_check_ref(storage.rhoe_snap)
1640 .get(id_patch)
1641 .get_buf()
1642 .complete_event_state(resulting_events);
1643 shambase::get_check_ref(storage.rho_vel_snap)
1644 .get(id_patch)
1645 .get_buf()
1646 .complete_event_state(resulting_events);
1647 }
1648
1649 void apply_refine_new(
1650 u32 cur_idx,
1651 BlockCoord cur_coords,
1652 std::array<u32, 8> new_blocks,
1653 std::array<BlockCoord, 8> new_block_coords,
1654 RefineCellAccessor acc) const {
1655
1656 auto get_coord_ref = [](u32 i) -> std::array<u32, dim> {
1657 constexpr u32 NsideBlockPow = 1;
1658 constexpr u32 Nside = 1U << NsideBlockPow;
1659
1660 if constexpr (dim == 3) {
1661 const u32 tmp = i >> NsideBlockPow;
1662 return {i % Nside, (tmp) % Nside, (tmp) >> NsideBlockPow};
1663 }
1664 };
1665
1666 auto get_index_block = [](std::array<u32, dim> coord) -> u32 {
1667 constexpr u32 NsideBlockPow = 1;
1668 constexpr u32 Nside = 1U << NsideBlockPow;
1669
1670 if constexpr (dim == 3) {
1671 return coord[0] + Nside * coord[1] + Nside * Nside * coord[2];
1672 }
1673 };
1674
1675 auto get_gid_write = [&](std::array<u32, dim> &glid) -> u32 {
1676 // First, get the block id (it's the block to be refine) in wich the new cell glid
1677 // is located.
1678 std::array<u32, dim> bid
1679 = {glid[0] >> AMRBlock::NsideBlockPow,
1680 glid[1] >> AMRBlock::NsideBlockPow,
1681 glid[2] >> AMRBlock::NsideBlockPow};
1682
1683 // get the new global block id
1684 auto new_glob_id = new_blocks[get_index_block(bid)] * AMRBlock::block_size;
1685
1686 // then added to new_glob_id the local index (between 0 and 7) of the generated
1687 // cells to get. This give the global ids of the new generated cells.
1688 return new_glob_id
1689 + AMRBlock::get_index(
1690 {glid[0] % AMRBlock::Nside,
1691 glid[1] % AMRBlock::Nside,
1692 glid[2] % AMRBlock::Nside});
1693 };
1694
1695 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
1696
1697 auto [lx, ly, lz] = get_coord_ref(loc_id);
1698 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
1699
1700 // // cell size in the refined block
1701 Tscal delta_cell = cell_sizes[cur_idx];
1702 Tscal c_offset = delta_cell * 0.25;
1703 std::array<f64_3, AMRBlock::block_size> child_center_offsets;
1704 child_center_offsets[0] = {-c_offset, -c_offset, -c_offset}; /*(0,0,0) */
1705 child_center_offsets[1] = {c_offset, -c_offset, -c_offset}; /*(1,0,0)*/
1706 child_center_offsets[2] = {-c_offset, c_offset, -c_offset}; /* (0,1,0)*/
1707 child_center_offsets[3] = {c_offset, c_offset, -c_offset}; /*(1,1,0)*/
1708 child_center_offsets[4] = {-c_offset, -c_offset, c_offset}; /*(0,0,1)*/
1709 child_center_offsets[5] = {c_offset, -c_offset, c_offset}; /*(1,0,1)*/
1710 child_center_offsets[6] = {-c_offset, c_offset, c_offset}; /*(0,1,1)*/
1711 child_center_offsets[7] = {c_offset, c_offset, c_offset}; /*(1,1,1)*/
1712
1713 auto cons_var_slopes = get_3d_grad_cons<Tvec, Minmod>(
1714 cell_sizes,
1715 AMRBlock::block_size,
1716 old_cell_idx,
1717 acc.cell_graph_xp,
1718 acc.cell_graph_xm,
1719 acc.cell_graph_yp,
1720 acc.cell_graph_ym,
1721 acc.cell_graph_zp,
1722 acc.cell_graph_zm,
1723 [=](u32 id) {
1724 return acc.rho_old_snap[id];
1725 },
1726 [=](u32 id) {
1727 return acc.rho_vel_old_snap[id];
1728 },
1729 [=](u32 id) {
1730 return acc.rhoE_old_snap[id];
1731 });
1732
1733 std::array<f64, AMRBlock::block_size> _rho_block;
1734 std::array<f64_3, AMRBlock::block_size> _rho_vel_block;
1735 std::array<f64, AMRBlock::block_size> _rhoE_block;
1736
1737 int mul_second_order
1738 = (acc.amr_ref_interp_mode == AMRInterpMode::SECOND_ORDER) ? 1 : 0;
1739
1740 bool do_second_order = true;
1741
1742 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
1743 shammath::ConsState<Tvec> cons_var_interp
1744 = child_center_offsets[subdiv_lid][0] * cons_var_slopes[0]
1745 + child_center_offsets[subdiv_lid][1] * cons_var_slopes[1]
1746 + child_center_offsets[subdiv_lid][2] * cons_var_slopes[2];
1747
1748 _rho_block[subdiv_lid]
1749 = acc.rho_old_snap[old_cell_idx] + cons_var_interp.rho * mul_second_order;
1750 _rho_vel_block[subdiv_lid] = acc.rho_vel_old_snap[old_cell_idx]
1751 + cons_var_interp.rhovel * mul_second_order;
1752 _rhoE_block[subdiv_lid]
1753 = acc.rhoE_old_snap[old_cell_idx] + cons_var_interp.rhoe * mul_second_order;
1754
1755 const auto e_int
1756 = _rhoE_block[subdiv_lid]
1757 - 0.5 * sham::dot(_rho_vel_block[subdiv_lid], _rho_vel_block[subdiv_lid])
1758 / (_rho_block[subdiv_lid]);
1759 do_second_order
1760 = do_second_order && (_rho_block[subdiv_lid] > 0.0) && (e_int > 0.0);
1761 }
1762
1763 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
1764
1765 auto [sx, sy, sz] = get_coord_ref(subdiv_lid);
1766
1767 std::array<u32, 3> glid = {lx * 2 + sx, ly * 2 + sy, lz * 2 + sz};
1768
1769 u32 new_cell_idx = get_gid_write(glid);
1770
1771 acc.rho[new_cell_idx] = (do_second_order) ? _rho_block[subdiv_lid]
1772 : acc.rho_old_snap[old_cell_idx];
1773 acc.rho_vel[new_cell_idx] = (do_second_order)
1774 ? _rho_vel_block[subdiv_lid]
1775 : acc.rho_vel_old_snap[old_cell_idx];
1776 acc.rhoE[new_cell_idx] = (do_second_order) ? _rhoE_block[subdiv_lid]
1777 : acc.rhoE_old_snap[old_cell_idx];
1778 }
1779 }
1780 }
1781
1782 void apply_derefine_new(
1783 std::array<u32, 8> old_blocks,
1784 std::array<BlockCoord, 8> old_coords,
1785 u32 new_cell,
1786 BlockCoord new_coord,
1787
1788 RefineCellAccessor acc) const {
1789
1790 std::array<f64, AMRBlock::block_size> rho_block;
1791 std::array<f64_3, AMRBlock::block_size> rho_vel_block;
1792 std::array<f64, AMRBlock::block_size> rhoE_block;
1793
1794 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1795 rho_block[cell_id] = {};
1796 rho_vel_block[cell_id] = {};
1797 rhoE_block[cell_id] = {};
1798 }
1799
1800 // for each siblings block, perform restriction from its 8 children cells
1801 for (u32 pid = 0; pid < 8; pid++) {
1802 auto rho_pid = rho_block[pid];
1803 auto rho_vel_pid = rho_vel_block[pid];
1804 auto rhoe_pid = rhoE_block[pid];
1805
1806 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1807 rho_pid += acc.rho[old_blocks[pid] * AMRBlock::block_size + cell_id];
1808 rho_vel_pid += acc.rho_vel[old_blocks[pid] * AMRBlock::block_size + cell_id];
1809 rhoe_pid += acc.rhoE[old_blocks[pid] * AMRBlock::block_size + cell_id];
1810 }
1811 rho_block[pid] = rho_pid * (1. / 8.);
1812 rho_vel_block[pid] = rho_vel_pid * (1. / 8.);
1813 rhoE_block[pid] = rhoe_pid * (1. / 8.);
1814 }
1815
1816 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
1817 u32 newcell_idx = new_cell * AMRBlock::block_size + cell_id;
1818 acc.rho[newcell_idx] = rho_block[cell_id];
1819 acc.rho_vel[newcell_idx] = rho_vel_block[cell_id];
1820 acc.rhoE[newcell_idx] = rhoE_block[cell_id];
1821 }
1822 }
1823 };
1824
1828 class RefineCritPseudoGradientAccessor {
1829 public:
1830 Tscal one_over_Nside = 1. / AMRBlock::Nside;
1831 const TgridVec *block_low_bound;
1832 const TgridVec *block_high_bound;
1833 const Tscal *block_rho;
1834 const f64 *block_pressure;
1835 const f64_3 *block_velocity;
1836
1837 const Tscal *rho_cons;
1838 Tscal error_min;
1839 Tscal error_max;
1840 u32 nblock_per_patch;
1841
1842 AMRGraphLinkiterator cell_graph_xp;
1843 AMRGraphLinkiterator cell_graph_xm;
1844 AMRGraphLinkiterator cell_graph_yp;
1845 AMRGraphLinkiterator cell_graph_ym;
1846 AMRGraphLinkiterator cell_graph_zp;
1847 AMRGraphLinkiterator cell_graph_zm;
1848
1849 RefineCritPseudoGradientAccessor(
1850 sham::EventList &depends_list,
1851 Storage &storage,
1852 u64 id_patch,
1855 Tscal err_min,
1856 Tscal err_max)
1857 : error_min(err_min), error_max(err_max),
1858 cell_graph_xp(
1859 shambase::get_check_ref(storage.cell_graph_edge)
1860 .get_refs_dir(Direction_::xp)
1861 .get(id_patch)
1862 .get()
1863 .get_read_access(depends_list)),
1864 cell_graph_xm(
1865 shambase::get_check_ref(storage.cell_graph_edge)
1866 .get_refs_dir(Direction_::xm)
1867 .get(id_patch)
1868 .get()
1869 .get_read_access(depends_list)),
1870 cell_graph_yp(
1871 shambase::get_check_ref(storage.cell_graph_edge)
1872 .get_refs_dir(Direction_::yp)
1873 .get(id_patch)
1874 .get()
1875 .get_read_access(depends_list)),
1876 cell_graph_ym(
1877 shambase::get_check_ref(storage.cell_graph_edge)
1878 .get_refs_dir(Direction_::ym)
1879 .get(id_patch)
1880 .get()
1881 .get_read_access(depends_list)),
1882 cell_graph_zp(
1883 shambase::get_check_ref(storage.cell_graph_edge)
1884 .get_refs_dir(Direction_::zp)
1885 .get(id_patch)
1886 .get()
1887 .get_read_access(depends_list)),
1888 cell_graph_zm(
1889 shambase::get_check_ref(storage.cell_graph_edge)
1890 .get_refs_dir(Direction_::zm)
1891 .get(id_patch)
1892 .get()
1893 .get_read_access(depends_list))
1894
1895 {
1896 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
1897 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
1898 rho_cons = pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1899 .get_buf()
1900 .get_read_access(depends_list);
1901
1902 nblock_per_patch = pdat.get_obj_cnt();
1903
1904 block_rho = shambase::get_check_ref(storage.rho_primitive)
1905 .get_buf(id_patch)
1906 .get_read_access(depends_list);
1907 block_pressure = shambase::get_check_ref(storage.press)
1908 .get_buf(id_patch)
1909 .get_read_access(depends_list);
1910 block_velocity = shambase::get_check_ref(storage.vel)
1911 .get_buf(id_patch)
1912 .get_read_access(depends_list);
1913 }
1914
1915 void finalize_new(
1916 sham::EventList &resulting_events,
1917 Storage &storage,
1918 u64 id_patch,
1921 Tscal err_min,
1922 Tscal err_max) {
1923
1924 pdat.get_field<i64_3>(0).get_buf().complete_event_state(resulting_events);
1925 pdat.get_field<i64_3>(1).get_buf().complete_event_state(resulting_events);
1926 pdat.get_field<Tscal>(pdat.pdl().get_field_idx<Tscal>("rho"))
1927 .get_buf()
1928 .complete_event_state(resulting_events);
1929
1930 shambase::get_check_ref(storage.cell_graph_edge)
1931 .get_refs_dir(Direction_::xp)
1932 .get(id_patch)
1933 .get()
1934 .complete_event_state(resulting_events);
1935
1936 shambase::get_check_ref(storage.cell_graph_edge)
1937 .get_refs_dir(Direction_::xm)
1938 .get(id_patch)
1939 .get()
1940 .complete_event_state(resulting_events);
1941
1942 shambase::get_check_ref(storage.cell_graph_edge)
1943 .get_refs_dir(Direction_::yp)
1944 .get(id_patch)
1945 .get()
1946 .complete_event_state(resulting_events);
1947 shambase::get_check_ref(storage.cell_graph_edge)
1948 .get_refs_dir(Direction_::ym)
1949 .get(id_patch)
1950 .get()
1951 .complete_event_state(resulting_events);
1952 shambase::get_check_ref(storage.cell_graph_edge)
1953 .get_refs_dir(Direction_::zp)
1954 .get(id_patch)
1955 .get()
1956 .complete_event_state(resulting_events);
1957 shambase::get_check_ref(storage.cell_graph_edge)
1958 .get_refs_dir(Direction_::zm)
1959 .get(id_patch)
1960 .get()
1961 .complete_event_state(resulting_events);
1962
1963 shambase::get_check_ref(storage.rho_primitive)
1964 .get_buf(id_patch)
1965 .complete_event_state(resulting_events);
1966
1967 shambase::get_check_ref(storage.press)
1968 .get_buf(id_patch)
1969 .complete_event_state(resulting_events);
1970
1971 shambase::get_check_ref(storage.vel)
1972 .get_buf(id_patch)
1973 .complete_event_state(resulting_events);
1974 }
1975
1976 void refine_criterion_new(
1977 u32 block_id,
1978 RefineCritPseudoGradientAccessor acc,
1979 bool &should_refine,
1980 bool &should_derefine) const {
1981 TgridVec low_bound = acc.block_low_bound[block_id];
1982 TgridVec high_bound = acc.block_high_bound[block_id];
1983
1984 // Tscal block_rho_slope = shambase::VectorProperties<Tscal>::get_zero();
1985 // for (u32 i = 0; i < AMRBlock::block_size; i++) {
1986 // block_rho_slope = sham::details::g_sycl_max(
1987 // block_rho_slope,
1988 // modif_second_derivative<Tscal, Tvec>(
1989 // i + block_id * AMRBlock::block_size,
1990 // cell_graph_xp,
1991 // cell_graph_xm,
1992 // cell_graph_yp,
1993 // cell_graph_ym,
1994 // cell_graph_zp,
1995 // cell_graph_zm,
1996 // [=](u32 id) {
1997 // return acc.block_rho[id];
1998 // }));
1999 // }
2000
2001 Tscal block_rho_slope = shambase::VectorProperties<Tscal>::get_zero();
2002 // for (u32 i = 0; i < AMRBlock::block_size; i++) {
2003 // block_rho_slope = sham::details::g_sycl_max(
2004 // block_rho_slope,
2005 // // baryonic_normalized_slope_criterion<Tscal>
2006 // // get_pseudo_grad<Tscal, Tvec>
2007 // (
2008 // i + block_id * AMRBlock::block_size,
2009 // cell_graph_xp,
2010 // cell_graph_xm,
2011 // cell_graph_yp,
2012 // cell_graph_ym,
2013 // cell_graph_zp,
2014 // cell_graph_zm,
2015 // [=](u32 id) {
2016 // // return rho_cons[id];
2017 // return acc.block_rho[id];
2018 // }));
2019 // }
2020
2021 Tscal block_press_grad = shambase::VectorProperties<Tscal>::get_zero();
2022 for (u32 i = 0; i < AMRBlock::block_size; i++) {
2023 block_press_grad = sham::details::g_sycl_max(
2024 block_press_grad,
2025 get_pseudo_grad<Tscal, Tvec>(
2026 i + block_id * AMRBlock::block_size,
2027 cell_graph_xp,
2028 cell_graph_xm,
2029 cell_graph_yp,
2030 cell_graph_ym,
2031 cell_graph_zp,
2032 cell_graph_zm,
2033 [=, this](u32 id) {
2034 return block_pressure[id];
2035 }));
2036 }
2037
2038 Tscal error = sham::details::g_sycl_max(
2039 block_press_grad, sham::details::g_sycl_max(block_rho_slope, 0.0));
2040
2041 should_refine = false;
2042 should_derefine = false;
2043 if (error > error_max) {
2044 should_refine = true;
2045 } else if (error < (error_min * error_max)) {
2046 should_derefine = true;
2047 }
2048
2049 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
2050 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
2051 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
2052 }
2053 };
2054
2058 class RefineCritShearAccessor {
2059 public:
2060 Tscal one_over_Nside = 1. / AMRBlock::Nside;
2061 const TgridVec *block_low_bound;
2062 const TgridVec *block_high_bound;
2063 const Tscal *block_rho;
2064 const f64 *block_pressure;
2065 const f64_3 *block_velocity;
2066
2067 Tscal threshold;
2068 Tscal gamma;
2069 Tscal dxfact;
2070
2071 AMRGraphLinkiterator cell_graph_xp;
2072 AMRGraphLinkiterator cell_graph_xm;
2073 AMRGraphLinkiterator cell_graph_yp;
2074 AMRGraphLinkiterator cell_graph_ym;
2075 AMRGraphLinkiterator cell_graph_zp;
2076 AMRGraphLinkiterator cell_graph_zm;
2077
2078 RefineCritShearAccessor(
2079 sham::EventList &depends_list,
2080 Storage &storage,
2081 u64 id_patch,
2084 Tscal threshold,
2085 Tscal gamma,
2086 Tscal dxfact)
2087 : threshold(threshold), gamma(gamma), dxfact(dxfact),
2088 cell_graph_xp(
2089 shambase::get_check_ref(storage.cell_graph_edge)
2090 .get_refs_dir(Direction_::xp)
2091 .get(id_patch)
2092 .get()
2093 .get_read_access(depends_list)),
2094 cell_graph_xm(
2095 shambase::get_check_ref(storage.cell_graph_edge)
2096 .get_refs_dir(Direction_::xm)
2097 .get(id_patch)
2098 .get()
2099 .get_read_access(depends_list)),
2100 cell_graph_yp(
2101 shambase::get_check_ref(storage.cell_graph_edge)
2102 .get_refs_dir(Direction_::yp)
2103 .get(id_patch)
2104 .get()
2105 .get_read_access(depends_list)),
2106 cell_graph_ym(
2107 shambase::get_check_ref(storage.cell_graph_edge)
2108 .get_refs_dir(Direction_::ym)
2109 .get(id_patch)
2110 .get()
2111 .get_read_access(depends_list)),
2112 cell_graph_zp(
2113 shambase::get_check_ref(storage.cell_graph_edge)
2114 .get_refs_dir(Direction_::zp)
2115 .get(id_patch)
2116 .get()
2117 .get_read_access(depends_list)),
2118 cell_graph_zm(
2119 shambase::get_check_ref(storage.cell_graph_edge)
2120 .get_refs_dir(Direction_::zm)
2121 .get(id_patch)
2122 .get()
2123 .get_read_access(depends_list))
2124
2125 {
2126 block_low_bound = pdat.get_field<TgridVec>(0).get_buf().get_read_access(depends_list);
2127 block_high_bound = pdat.get_field<TgridVec>(1).get_buf().get_read_access(depends_list);
2128
2129 block_rho = shambase::get_check_ref(storage.rho_primitive)
2130 .get_buf(id_patch)
2131 .get_read_access(depends_list);
2132 block_pressure = shambase::get_check_ref(storage.press)
2133 .get_buf(id_patch)
2134 .get_read_access(depends_list);
2135 block_velocity = shambase::get_check_ref(storage.vel)
2136 .get_buf(id_patch)
2137 .get_read_access(depends_list);
2138 }
2139
2140 void finalize_new(
2141 sham::EventList &resulting_events,
2142 Storage &storage,
2143 u64 id_patch,
2146 Tscal threshold,
2147 Tscal gamma,
2148 Tscal dxfact) {
2149
2150 pdat.get_field<i64_3>(0).get_buf().complete_event_state(resulting_events);
2151 pdat.get_field<i64_3>(1).get_buf().complete_event_state(resulting_events);
2152
2153 shambase::get_check_ref(storage.cell_graph_edge)
2154 .get_refs_dir(Direction_::xp)
2155 .get(id_patch)
2156 .get()
2157 .complete_event_state(resulting_events);
2158
2159 shambase::get_check_ref(storage.cell_graph_edge)
2160 .get_refs_dir(Direction_::xm)
2161 .get(id_patch)
2162 .get()
2163 .complete_event_state(resulting_events);
2164
2165 shambase::get_check_ref(storage.cell_graph_edge)
2166 .get_refs_dir(Direction_::yp)
2167 .get(id_patch)
2168 .get()
2169 .complete_event_state(resulting_events);
2170 shambase::get_check_ref(storage.cell_graph_edge)
2171 .get_refs_dir(Direction_::ym)
2172 .get(id_patch)
2173 .get()
2174 .complete_event_state(resulting_events);
2175 shambase::get_check_ref(storage.cell_graph_edge)
2176 .get_refs_dir(Direction_::zp)
2177 .get(id_patch)
2178 .get()
2179 .complete_event_state(resulting_events);
2180 shambase::get_check_ref(storage.cell_graph_edge)
2181 .get_refs_dir(Direction_::zm)
2182 .get(id_patch)
2183 .get()
2184 .complete_event_state(resulting_events);
2185
2186 shambase::get_check_ref(storage.rho_primitive)
2187 .get_buf(id_patch)
2188 .complete_event_state(resulting_events);
2189
2190 shambase::get_check_ref(storage.press)
2191 .get_buf(id_patch)
2192 .complete_event_state(resulting_events);
2193
2194 shambase::get_check_ref(storage.vel)
2195 .get_buf(id_patch)
2196 .complete_event_state(resulting_events);
2197 }
2198
2199 void refine_criterion_new(
2200 u32 block_id,
2201 RefineCritShearAccessor acc,
2202 bool &should_refine,
2203 bool &should_derefine) const {
2204 TgridVec low_bound = acc.block_low_bound[block_id];
2205 TgridVec high_bound = acc.block_high_bound[block_id];
2206
2207 Tvec lower_flt = low_bound.template convert<Tscal>() * dxfact;
2208 Tvec upper_flt = high_bound.template convert<Tscal>() * dxfact;
2209
2210 Tvec block_cell_size = (upper_flt - lower_flt) * one_over_Nside;
2211
2212 Tscal block_normalized_shear = shambase::VectorProperties<Tscal>::get_zero();
2213 for (u32 i = 0; i < AMRBlock::block_size; i++) {
2214 auto cell_id = i + block_id * AMRBlock::block_size;
2215 auto cs = sycl::sqrt(gamma * acc.block_pressure[cell_id] / acc.block_rho[cell_id]);
2216 block_normalized_shear = sham::details::g_sycl_max(
2217 block_normalized_shear,
2218 normalized_shear<Tvec>(
2219 cell_id,
2220 cs,
2221 block_cell_size,
2222
2223 cell_graph_xp,
2224 cell_graph_xm,
2225 cell_graph_yp,
2226 cell_graph_ym,
2227 cell_graph_zp,
2228 cell_graph_zm,
2229 [=](u32 id) {
2230 return acc.block_velocity[id];
2231 }));
2232 }
2233 should_refine = false;
2234 should_derefine = false;
2235 if (block_normalized_shear > threshold * threshold) {
2236 should_refine = true;
2237 } else if (block_normalized_shear < 0.25 * threshold * threshold) {
2238 should_derefine = true;
2239 }
2240
2241 should_refine = should_refine && (high_bound.x() - low_bound.x() > AMRBlock::Nside);
2242 should_refine = should_refine && (high_bound.y() - low_bound.y() > AMRBlock::Nside);
2243 should_refine = should_refine && (high_bound.z() - low_bound.z() > AMRBlock::Nside);
2244 }
2245 };
2246
2250 class RefineCellAccessorAutogravity {
2251 public:
2252 f64 *rho;
2253 f64_3 *rho_vel;
2254 f64 *rhoE;
2255 f64 *phi_old;
2256 f64 *phi_new;
2257
2258 u64 p_id;
2259 // f64* cell_sizes;
2260
2261 // this will be needed for interpolation during refinement
2262 AMRGraphLinkiterator cell_graph_xp;
2263 AMRGraphLinkiterator cell_graph_xm;
2264 AMRGraphLinkiterator cell_graph_yp;
2265 AMRGraphLinkiterator cell_graph_ym;
2266 AMRGraphLinkiterator cell_graph_zp;
2267 AMRGraphLinkiterator cell_graph_zm;
2268
2269 RefineCellAccessorAutogravity(
2270 sham::EventList &depends_list,
2271 Storage &storage,
2272 u64 &id_patch,
2274 : cell_graph_xp(
2275 shambase::get_check_ref(storage.cell_graph_edge)
2276 .get_refs_dir(Direction::xp)
2277 .get(id_patch)
2278 .get()
2279 .get_read_access(depends_list)),
2280 cell_graph_xm(
2281 shambase::get_check_ref(storage.cell_graph_edge)
2282 .get_refs_dir(Direction::xm)
2283 .get(id_patch)
2284 .get()
2285 .get_read_access(depends_list)),
2286 cell_graph_yp(
2287 shambase::get_check_ref(storage.cell_graph_edge)
2288 .get_refs_dir(Direction::yp)
2289 .get(id_patch)
2290 .get()
2291 .get_read_access(depends_list)),
2292 cell_graph_ym(
2293 shambase::get_check_ref(storage.cell_graph_edge)
2294 .get_refs_dir(Direction::ym)
2295 .get(id_patch)
2296 .get()
2297 .get_read_access(depends_list)),
2298 cell_graph_zp(
2299 shambase::get_check_ref(storage.cell_graph_edge)
2300 .get_refs_dir(Direction::zp)
2301 .get(id_patch)
2302 .get()
2303 .get_read_access(depends_list)),
2304 cell_graph_zm(
2305 shambase::get_check_ref(storage.cell_graph_edge)
2306 .get_refs_dir(Direction::zm)
2307 .get(id_patch)
2308 .get()
2309 .get_read_access(depends_list))
2310
2311 {
2312 p_id = id_patch;
2313 rho = pdat.get_field<f64>(2).get_buf().get_write_access(depends_list);
2314 rho_vel = pdat.get_field<f64_3>(3).get_buf().get_write_access(depends_list);
2315 rhoE = pdat.get_field<f64>(4).get_buf().get_write_access(depends_list);
2316 phi_old = pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi_old"))
2317 .get_buf()
2318 .get_write_access(depends_list);
2319 phi_new = pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi"))
2320 .get_buf()
2321 .get_write_access(depends_list);
2322 }
2323
2324 void finalize_new(
2325 sham::EventList &resulting_events,
2326 Storage &storage,
2327 u64 &id_patch,
2329 pdat.get_field<f64>(2).get_buf().complete_event_state(resulting_events);
2330 pdat.get_field<f64_3>(3).get_buf().complete_event_state(resulting_events);
2331 pdat.get_field<f64>(4).get_buf().complete_event_state(resulting_events);
2332 pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi_old"))
2333 .get_buf()
2334 .complete_event_state(resulting_events);
2335 pdat.get_field<f64>(pdat.pdl().get_field_idx<Tscal>("phi"))
2336 .get_buf()
2337 .complete_event_state(resulting_events);
2338
2339 shambase::get_check_ref(storage.cell_graph_edge)
2340 .get_refs_dir(Direction_::xp)
2341 .get(id_patch)
2342 .get()
2343 .complete_event_state(resulting_events);
2344 shambase::get_check_ref(storage.cell_graph_edge)
2345 .get_refs_dir(Direction_::xm)
2346 .get(id_patch)
2347 .get()
2348 .complete_event_state(resulting_events);
2349 shambase::get_check_ref(storage.cell_graph_edge)
2350 .get_refs_dir(Direction_::yp)
2351 .get(id_patch)
2352 .get()
2353 .complete_event_state(resulting_events);
2354 shambase::get_check_ref(storage.cell_graph_edge)
2355 .get_refs_dir(Direction_::ym)
2356 .get(id_patch)
2357 .get()
2358 .complete_event_state(resulting_events);
2359 shambase::get_check_ref(storage.cell_graph_edge)
2360 .get_refs_dir(Direction_::zp)
2361 .get(id_patch)
2362 .get()
2363 .complete_event_state(resulting_events);
2364 shambase::get_check_ref(storage.cell_graph_edge)
2365 .get_refs_dir(Direction_::zm)
2366 .get(id_patch)
2367 .get()
2368 .complete_event_state(resulting_events);
2369 }
2370
2371 void apply_refine_new(
2372 u32 cur_idx,
2373 BlockCoord cur_coords,
2374 std::array<u32, 8> new_blocks,
2375 std::array<BlockCoord, 8> new_block_coords,
2376 RefineCellAccessorAutogravity acc) const {
2377
2378 auto get_coord_ref = [](u32 i) -> std::array<u32, dim> {
2379 constexpr u32 NsideBlockPow = 1;
2380 constexpr u32 Nside = 1U << NsideBlockPow;
2381
2382 if constexpr (dim == 3) {
2383 const u32 tmp = i >> NsideBlockPow;
2384 return {i % Nside, (tmp) % Nside, (tmp) >> NsideBlockPow};
2385 }
2386 };
2387
2388 auto get_index_block = [](std::array<u32, dim> coord) -> u32 {
2389 constexpr u32 NsideBlockPow = 1;
2390 constexpr u32 Nside = 1U << NsideBlockPow;
2391
2392 if constexpr (dim == 3) {
2393 return coord[0] + Nside * coord[1] + Nside * Nside * coord[2];
2394 }
2395 };
2396
2397 auto get_gid_write = [&](std::array<u32, dim> &glid) -> u32 {
2398 // First, get the block id (it's the block to be refine) in wich the new cell glid
2399 // is located.
2400 std::array<u32, dim> bid
2401 = {glid[0] >> AMRBlock::NsideBlockPow,
2402 glid[1] >> AMRBlock::NsideBlockPow,
2403 glid[2] >> AMRBlock::NsideBlockPow};
2404
2405 // get the new global block id
2406 auto new_glob_id = new_blocks[get_index_block(bid)] * AMRBlock::block_size;
2407
2408 // then added to new_glob_id the local index (between 0 and 7) of the generated
2409 // cells to get. This give the global ids of the new generated cells.
2410 return new_glob_id
2411 + AMRBlock::get_index(
2412 {glid[0] % AMRBlock::Nside,
2413 glid[1] % AMRBlock::Nside,
2414 glid[2] % AMRBlock::Nside});
2415 };
2416
2417 std::array<f64, AMRBlock::block_size> old_rho_block;
2418 std::array<f64_3, AMRBlock::block_size> old_rho_vel_block;
2419 std::array<f64, AMRBlock::block_size> old_rhoE_block;
2420 std::array<f64, AMRBlock::block_size> old_phi_old_block;
2421 std::array<f64, AMRBlock::block_size> old_phi_new_block;
2422
2423 // save old block
2424 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
2425
2426 auto [lx, ly, lz] = get_coord_ref(loc_id);
2427 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
2428 old_rho_block[loc_id] = acc.rho[old_cell_idx];
2429 old_rho_vel_block[loc_id] = acc.rho_vel[old_cell_idx];
2430 old_rhoE_block[loc_id] = acc.rhoE[old_cell_idx];
2431 old_phi_old_block[loc_id] = acc.phi_old[old_cell_idx];
2432 old_phi_new_block[loc_id] = acc.phi_new[old_cell_idx];
2433 }
2434
2435 for (u32 loc_id = 0; loc_id < AMRBlock::block_size; loc_id++) {
2436
2437 auto [lx, ly, lz] = get_coord_ref(loc_id);
2438 u32 old_cell_idx = cur_idx * AMRBlock::block_size + loc_id;
2439
2440 // // // cell size in the refined block
2441 // // Tscal delta_cell = cell_sizes[cur_idx];
2442 // // Tscal c_offset = delta_cell * 0.25;
2443 // // std::array<f64_3, AMRBlock::block_size> child_center_offsets;
2444 // // child_center_offsets[0] = {-c_offset, -c_offset, -c_offset}; /*(0,0,0) */
2445 // // child_center_offsets[1] = {c_offset, -c_offset, -c_offset}; /*(1,0,0)*/
2446 // // child_center_offsets[2] = {-c_offset, c_offset, -c_offset}; /* (0,1,0)*/
2447 // // child_center_offsets[3] = {c_offset, c_offset, -c_offset}; /*(1,1,0)*/
2448 // // child_center_offsets[4] = {-c_offset, -c_offset, c_offset}; /*(0,0,1)*/
2449 // // child_center_offsets[5] = {c_offset, -c_offset, c_offset}; /*(1,0,1)*/
2450 // // child_center_offsets[6] = {-c_offset, c_offset, c_offset}; /*(0,1,1)*/
2451 // // child_center_offsets[7] = {c_offset, c_offset, c_offset}; /*(1,1,1)*/
2452
2453 // auto cons_var_slopes = get_3d_grad_cons<Tvec, Minmod>(
2454 // old_cell_idx,
2455 // delta_cell,
2456 // cell_graph_xp,
2457 // cell_graph_xm,
2458 // cell_graph_yp,
2459 // cell_graph_ym,
2460 // cell_graph_zp,
2461 // cell_graph_zm,
2462 // [=](u32 id){
2463 // return acc.rho[id];
2464 // },
2465 // [=](u32 id){
2466 // return acc.rho_vel[id];
2467 // },
2468 // [=](u32 id){
2469 // return acc.rhoE[id];
2470 // });
2471
2472 Tscal rho_block = old_rho_block[loc_id];
2473 Tvec rho_vel_block = old_rho_vel_block[loc_id];
2474 Tscal rhoE_block = old_rhoE_block[loc_id];
2475 Tscal phi_old_block = old_phi_old_block[loc_id];
2476 Tscal phi_new_block = old_phi_new_block[loc_id];
2477
2478 for (u32 subdiv_lid = 0; subdiv_lid < 8; subdiv_lid++) {
2479
2480 auto [sx, sy, sz] = get_coord_ref(subdiv_lid);
2481
2482 std::array<u32, 3> glid = {lx * 2 + sx, ly * 2 + sy, lz * 2 + sz};
2483
2484 u32 new_cell_idx = get_gid_write(glid);
2485
2486 // shammath::ConsState<Tvec> cons_var_interp =
2487 // child_center_offsets[subdiv_lid][0] * cons_var_slopes[0] +
2488 // child_center_offsets[subdiv_lid][1] * cons_var_slopes[1] +
2489 // child_center_offsets[subdiv_lid][2] * cons_var_slopes[2];
2490
2491 // // acc.rho[new_cell_idx] = rho_block + cons_var_interp.rho ;
2492 // // acc.rho_vel[new_cell_idx] = rho_vel_block + cons_var_interp.rhovel;
2493 // // acc.rhoE[new_cell_idx] = rhoE_block + cons_var_interp.rhoe;
2494
2495 acc.rho[new_cell_idx] = rho_block;
2496 acc.rho_vel[new_cell_idx] = rho_vel_block;
2497 acc.rhoE[new_cell_idx] = rhoE_block;
2498 acc.phi_old[new_cell_idx] = phi_old_block;
2499 acc.phi_new[new_cell_idx] = phi_new_block;
2500 }
2501 }
2502 }
2503
2504 void apply_derefine_new(
2505 std::array<u32, 8> old_blocks,
2506 std::array<BlockCoord, 8> old_coords,
2507 u32 new_cell,
2508 BlockCoord new_coord,
2509
2510 RefineCellAccessorAutogravity acc) const {
2511
2512 std::array<f64, AMRBlock::block_size> rho_block;
2513 std::array<f64_3, AMRBlock::block_size> rho_vel_block;
2514 std::array<f64, AMRBlock::block_size> rhoE_block;
2515 std::array<f64, AMRBlock::block_size> phi_old_block;
2516 std::array<f64, AMRBlock::block_size> phi_new_block;
2517
2518 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2519 rho_block[cell_id] = {};
2520 rho_vel_block[cell_id] = {};
2521 rhoE_block[cell_id] = {};
2522 phi_old_block[cell_id] = {};
2523 phi_new_block[cell_id] = {};
2524 }
2525
2526 // for each siblings block, perform restriction from its 8 children cells
2527 for (u32 pid = 0; pid < 8; pid++) {
2528 auto rho_pid = rho_block[pid];
2529 auto rho_vel_pid = rho_vel_block[pid];
2530 auto rhoe_pid = rhoE_block[pid];
2531 auto phi_old_pid = phi_old_block[pid];
2532 auto phi_new_pid = phi_new_block[pid];
2533
2534 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2535 rho_pid += acc.rho[old_blocks[pid] * AMRBlock::block_size + cell_id];
2536 rho_vel_pid += acc.rho_vel[old_blocks[pid] * AMRBlock::block_size + cell_id];
2537 rhoe_pid += acc.rhoE[old_blocks[pid] * AMRBlock::block_size + cell_id];
2538 phi_old_pid += acc.phi_old[old_blocks[pid] * AMRBlock::block_size + cell_id];
2539 phi_new_pid += acc.phi_new[old_blocks[pid] * AMRBlock::block_size + cell_id];
2540 }
2541 rho_block[pid] = rho_pid * (1. / 8.);
2542 rho_vel_block[pid] = rho_vel_pid * (1. / 8.);
2543 rhoE_block[pid] = rhoe_pid * (1. / 8.);
2544 // phi_old_block[pid] = phi_old_pid * (1. / 8.);
2545 // phi_new_block[pid] = phi_new_pid * (1. / 8.);
2546 }
2547
2548 for (u32 cell_id = 0; cell_id < AMRBlock::block_size; cell_id++) {
2549 u32 newcell_idx = new_cell * AMRBlock::block_size + cell_id;
2550 acc.rho[newcell_idx] = rho_block[cell_id];
2551 acc.rho_vel[newcell_idx] = rho_vel_block[cell_id];
2552 acc.rhoE[newcell_idx] = rhoE_block[cell_id];
2553
2554 // acc.phi_old[newcell_idx] = phi_old_block[cell_id];
2555 // acc.phi_new[newcell_idx] = phi_new_block[cell_id];
2556 }
2557 }
2558 };
2559
2560 using AMRmode_None = typename AMRMode<Tvec, TgridVec>::None;
2561 using AMRmode_DensityBased = typename AMRMode<Tvec, TgridVec>::DensityBased;
2562 using AMRmode_PseudoGradientBased = typename AMRMode<Tvec, TgridVec>::PseudoGradientBased;
2563 using AMRmode_JeansLengthBased = typename AMRMode<Tvec, TgridVec>::JeansLengthBased;
2564 using AMRmode_ShearBased = typename AMRMode<Tvec, TgridVec>::ShearBased;
2565
2566 bool has_cell_order_changed = false;
2567
2568 // get refine and derefine list
2571
2572 if (AMRmode_None *cfg = std::get_if<AMRmode_None>(&solver_config.amr_mode.config)) {
2573 // no refinment here turn around there is nothing to see
2574 } else {
2575 if (AMRmode_DensityBased *cfg
2576 = std::get_if<AMRmode_DensityBased>(&solver_config.amr_mode.config)) {
2577
2578 Tscal dxfact(solver_config.grid_coord_to_pos_fact);
2579 gen_refine_block_changes_new<RefineCritBlock>(
2580 refine_list, derefine_list, dxfact, cfg->crit_mass);
2581 }
2582
2583 else if (
2584 AMRmode_PseudoGradientBased *cfg
2585 = std::get_if<AMRmode_PseudoGradientBased>(&solver_config.amr_mode.config)) {
2586
2587 gen_refine_block_changes_new<RefineCritPseudoGradientAccessor>(
2588 refine_list, derefine_list, cfg->error_min, cfg->error_max);
2589 }
2590
2591 else if (
2592 AMRmode_ShearBased *cfg
2593 = std::get_if<AMRmode_ShearBased>(&solver_config.amr_mode.config)) {
2594 Tscal dxfact(solver_config.grid_coord_to_pos_fact);
2595 Tscal gamma(solver_config.eos_gamma);
2596
2597 gen_refine_block_changes_new<RefineCritShearAccessor>(
2598 refine_list, derefine_list, cfg->threshold, gamma, dxfact);
2599 }
2600
2602 enforce_two_to_one_refinement_new(std::move(refine_list));
2604 enforce_two_to_one_derefinement_new(std::move(derefine_list), std::move(refine_list));
2606 // Note that this only add new blocks at the end of the patchdata
2607 const AMRInterpMode amr_ref_interp_mode = solver_config.amr_interp_mode;
2608 bool change_refine = internal_refine_grid_new<RefineCellAccessor>(
2609 std::move(refine_list), amr_ref_interp_mode);
2610
2612 // Note that this will perform the merge then remove the old blocks
2613 // This is ok to call straight after the refine without edditing the index list in
2614 // derefine_list since no permutations were applied in internal_refine_grid_new and no cells
2615 // can be both refined and derefined in the same pass
2616 bool change_derefine = internal_derefine_grid_new<RefineCellAccessor>(
2617 std::move(derefine_list), amr_ref_interp_mode);
2618
2619 has_cell_order_changed = has_cell_order_changed || (change_refine || change_derefine);
2620
2621 if (has_cell_order_changed) {
2622 // Ensure that the blocks are sorted before refinement
2623 AMRSortBlocks block_sorter(context, solver_config, storage);
2624 block_sorter.reorder_amr_blocks();
2625 }
2626 }
2627}
2628
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:32
void add_event(sycl::event e)
Add an event to the list of events.
Definition EventList.hpp:88
Represents a collection of objects distributed across patches identified by a u64 id.
PatchDataLayer container class, the layout is described in patchdata_layout.
u32 get_obj_cnt() const
get the number of objects (particles) stored in this layer
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:318
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:67
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
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:112
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:89
void info_ln(std::string module_name, Types... var2)
Prints a log message with multiple arguments followed by a newline.
Definition logs.hpp:132
Patch object that contain generic patch information.
Definition Patch.hpp:33
u64 id_patch
unique key that identify the patch
Definition Patch.hpp:86