58 using Tscal = shambase::VecComponent<Tvec>;
59 static constexpr u32 dim = shambase::VectorProperties<Tvec>::dimension;
60 using Kernel = SPHKernel<Tscal>;
62 using Solver = Solver<Tvec, SPHKernel>;
63 using SolverConfig =
typename Solver::Config;
84 solver.solver_config.scheduler_conf.split_load_value = crit_split;
85 solver.solver_config.scheduler_conf.merge_load_value = crit_merge;
89 template<std::enable_if_t<dim == 3,
int> = 0>
90 inline Tvec get_box_dim_fcc_3d(Tscal dr,
u32 xcnt,
u32 ycnt,
u32 zcnt) {
91 return generic::setup::generators::get_box_dim(dr, xcnt, ycnt, zcnt);
94 inline void set_cfl_cour(Tscal cfl_cour) {
95 solver.solver_config.cfl_config.cfl_cour = cfl_cour;
97 inline void set_cfl_force(Tscal cfl_force) {
98 solver.solver_config.cfl_config.cfl_force = cfl_force;
100 inline void set_eta_sink(Tscal eta_sink) {
101 solver.solver_config.cfl_config.eta_sink = eta_sink;
104 inline Tscal get_time() {
return solver.get_time(); }
105 inline void set_time(Tscal t) { solver.set_time(t); }
106 inline Tscal get_dt_sph() {
return solver.get_dt_sph(); }
107 inline void set_next_dt(Tscal dt) { solver.set_next_dt(dt); }
108 inline Tscal get_cfl_multipler() {
return solver.get_cfl_multipler(); }
109 inline void set_cfl_multipler(Tscal lambda) { solver.set_cfl_multipler(lambda); }
111 inline void set_particle_mass(Tscal gpart_mass) {
112 solver.solver_config.gpart_mass = gpart_mass;
115 inline Tscal get_particle_mass() {
return solver.solver_config.gpart_mass; }
117 inline void resize_simulation_box(std::pair<Tvec, Tvec> box) {
118 ctx.set_coord_domain_bound({box.first, box.second});
121 SolverConfig gen_config_from_phantom_dump(PhantomDump &phdump,
bool bypass_error);
122 void init_from_phantom_dump(PhantomDump &phdump, Tscal hpart_fact_load = 1.0);
123 PhantomDump make_phantom_dump();
125 void do_vtk_dump(std::string filename,
bool add_patch_world_id) {
126 solver.vtk_do_dump(filename, add_patch_world_id);
129 void set_debug_dump(
bool _do_debug_dump, std::string _debug_dump_filename) {
130 solver.set_debug_dump(_do_debug_dump, _debug_dump_filename);
133 u64 get_total_part_count();
135 f64 total_mass_to_part_mass(
f64 totmass);
137 Tscal get_hfact() {
return Kernel::hfactd; }
139 Tscal rho_h(Tscal h) {
140 return shamrock::sph::rho_h(solver.solver_config.gpart_mass, h, Kernel::hfactd);
143 void add_cube_fcc_3d(Tscal dr, std::pair<Tvec, Tvec> _box);
144 void add_cube_hcp_3d(Tscal dr, std::pair<Tvec, Tvec> _box);
145 void add_cube_hcp_3d_v2(Tscal dr, std::pair<Tvec, Tvec> _box);
147 inline std::unique_ptr<modules::SPHSetup<Tvec, SPHKernel>> get_setup() {
148 return std::make_unique<modules::SPHSetup<Tvec, SPHKernel>>(
149 ctx, solver.solver_config, solver.storage);
171 void add_big_disc_3d(
183 inline void add_sink(Tscal mass, Tvec pos, Tvec velocity, Tscal accretion_radius) {
184 if (!ctx.is_scheduler_initialized()) {
186 "add_sink() requires that the scheduler has been initialized. "
187 "Call init_scheduler(...) before add_sink().");
192 "add_sink() requires that sink edges are registered. "
193 "Call init_scheduler(...) before add_sink().");
197 shamlog_debug_ln(
"SPH",
"add sink :", mass, pos, velocity, accretion_radius);
203 inline void set_field_value_lambda(
204 std::string field_name,
const std::function<T(Tvec)> pos_to_val,
const u32 offset) {
214 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
215 PatchDataField<Tvec> &
xyz = pdat.template get_field<Tvec>(ixyz);
216 PatchDataField<T> &f = pdat.template get_field<T>(ifield);
218 auto f_nvar = f.get_nvar();
219 if (offset >= f_nvar) {
221 "offset ({}) is out of bounds for field '{}' with nvar {}",
227 auto acc = f.get_buf().copy_to_stdvec();
228 auto acc_xyz =
xyz.get_buf().copy_to_stdvec();
230 u32 obj_cnt = pdat.get_obj_cnt();
231 for (
u32 i = 0; i < obj_cnt; i++) {
232 acc[i * f_nvar + offset] = pos_to_val(acc_xyz[i]);
235 f.get_buf().copy_from_stdvec(acc);
240 inline void overwrite_field_value(
241 std::string field_name,
242 const std::function<std::vector<T>(py::dict)> field_compute,
252 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
253 PatchDataField<T> &f = pdat.template get_field<T>(ifield);
255 auto f_nvar = f.get_nvar();
256 if (offset >= f_nvar) {
258 "offset ({}) is out of bounds for field '{}' with nvar {}",
264 auto result = field_compute(shamrock::pdat_to_dic(pdat));
266 if (result.size() != f.get_obj_cnt()) {
268 "result.size() != f.get_obj_cnt() ({} != {})",
273 auto acc = f.get_buf().copy_to_stdvec();
275 u32 obj_cnt = pdat.get_obj_cnt();
276 for (
u32 i = 0; i < obj_cnt; i++) {
277 acc[i * f_nvar + offset] = result[i];
280 f.get_buf().copy_from_stdvec(acc);
298 template<std::enable_if_t<dim == 3,
int> = 0>
310 Tscal G = solver.solver_config.get_constant_G();
313 using Config = SolverConfig;
314 using SolverConfigEOS =
typename Config::EOSConfig;
315 using SolverEOS_Adiabatic =
typename SolverConfigEOS::Adiabatic;
316 if (SolverEOS_Adiabatic *eos_config
317 = std::get_if<SolverEOS_Adiabatic>(&solver.solver_config.eos_config.config)) {
319 eos_gamma = eos_config->gamma;
329 auto sigma_profile = [=](Tscal r) {
331 constexpr Tscal sigma_0 = 1;
332 return sigma_0 * sycl::pow(r / r_in, -p);
335 auto cs_law = [=](Tscal r) {
336 return sycl::pow(r / r_in, -q);
339 auto rot_profile = [=](Tscal r) {
340 return sycl::sqrt(G * central_mass / r);
343 Tscal cs_in = H_r_in * rot_profile(r_in);
344 auto cs_profile = [&](Tscal r) {
345 return cs_law(r) * cs_in;
348 std::vector<Out> part_list;
355 return sigma_profile(r);
358 return cs_profile(r);
361 return rot_profile(r);
364 part_list.push_back(out);
367 Tscal part_mass = disc_mass / Npart;
369 using namespace shamrock::patch;
373 std::string log =
"";
380 std::vector<Tvec> vec_pos;
381 std::vector<Tvec> vec_vel;
382 std::vector<Tscal> vec_u;
383 std::vector<Tscal> vec_h;
385 std::vector<Tscal> vec_cs;
387 Tscal G = solver.solver_config.get_constant_G();
389 for (Out o : part_list) {
390 vec_pos.push_back(o.pos + center);
391 vec_vel.push_back(o.velocity);
397 vec_u.push_back(o.cs * o.cs / ( (eos_gamma - 1)));
398 vec_h.push_back(shamrock::sph::h_rho(part_mass, o.rho, Kernel::hfactd));
399 vec_cs.push_back(o.cs);
402 log += shambase::format(
403 "\n patch id={}, add N={} particles", ptch.
id_patch, vec_pos.size());
406 tmp.resize(vec_pos.size());
410 u32 len = vec_pos.size();
412 = tmp.get_field<Tvec>(sched.pdl_old().get_field_idx<Tvec>(
"xyz"));
413 sycl::buffer<Tvec> buf(vec_pos.data(), len);
414 f.override(buf, len);
418 u32 len = vec_pos.size();
420 = tmp.get_field<Tscal>(sched.pdl_old().get_field_idx<Tscal>(
"hpart"));
421 sycl::buffer<Tscal> buf(vec_h.data(), len);
422 f.override(buf, len);
426 u32 len = vec_pos.size();
428 = tmp.get_field<Tscal>(sched.pdl_old().get_field_idx<Tscal>(
"uint"));
429 sycl::buffer<Tscal> buf(vec_u.data(), len);
430 f.override(buf, len);
433 if (solver.solver_config.is_eos_locally_isothermal()) {
434 u32 len = vec_pos.size();
436 = tmp.get_field<Tscal>(sched.pdl_old().get_field_idx<Tscal>(
"soundspeed"));
437 sycl::buffer<Tscal> buf(vec_cs.data(), len);
438 f.override(buf, len);
442 u32 len = vec_pos.size();
444 = tmp.get_field<Tvec>(sched.pdl_old().get_field_idx<Tvec>(
"vxyz"));
445 sycl::buffer<Tvec> buf(vec_vel.data(), len);
446 f.override(buf, len);
449 pdat.insert_elements(tmp);
452 std::string log_gathered =
"";
460 ctx, solver.solver_config, solver.storage)
461 .update_load_balancing();
466 auto [m, M] = sched.get_box_tranform<Tvec>();
481 sched.check_patchdata_locality_correctness();
487 log += shambase::format(
488 "\n patch id={}, N={} particles", p.id_patch, pdat.get_obj_cnt());
499 template<std::enable_if_t<dim == 3,
int> = 0>
500 inline void add_cube_disc_3d(
513 using SolverConfigEOS =
typename Config::EOSConfig;
514 using SolverEOS_Adiabatic =
typename SolverConfigEOS::Adiabatic;
515 if (SolverEOS_Adiabatic *eos_config
516 = std::get_if<SolverEOS_Adiabatic>(&solver.solver_config.eos_config.config)) {
518 eos_gamma = eos_config->gamma;
524 auto cs = [&](Tscal u) {
525 return sycl::sqrt(eos_gamma * (eos_gamma - 1) * u);
528 auto U = [&](Tscal cs) {
529 return cs * cs / (eos_gamma * (eos_gamma - 1));
532 using namespace shamrock::patch;
536 std::string log =
"";
538 sched.for_each_local_patchdata([&](
const Patch &ptch, PatchDataLayer &pdat) {
541 shammath::CoordRange<Tvec> patch_coord = ptransf.to_obj_coord(ptch);
543 std::vector<Tvec> vec_acc;
544 std::vector<Tvec> vec_vel;
545 std::vector<Tscal> vec_u;
547 Tscal G = solver.solver_config.get_constant_G();
550 Npart, p, rho_0, m, r_in, r_out, q, [&](Tvec r, Tscal h) {
551 vec_acc.push_back(r + center);
553 Tscal R = sycl::length(r);
555 Tscal V = sycl::sqrt(G * cmass / R);
557 Tvec etheta = {-r.z(), 0, r.x()};
558 etheta /= sycl::length(etheta);
560 vec_vel.push_back(V * etheta);
563 Tscal cs = cs0 * sycl::pow(R, -q);
565 vec_u.push_back(U(cs));
568 log += shambase::format(
569 "\n patch id={}, add N={} particles", ptch.
id_patch, vec_acc.size());
571 PatchDataLayer tmp(sched.get_layout_ptr_old());
572 tmp.resize(vec_acc.size());
576 u32 len = vec_acc.size();
577 PatchDataField<Tvec> &f
578 = tmp.get_field<Tvec>(sched.pdl_old().get_field_idx<Tvec>(
"xyz"));
579 sycl::buffer<Tvec> buf(vec_acc.data(), len);
580 f.override(buf, len);
584 PatchDataField<Tscal> &f
585 = tmp.get_field<Tscal>(sched.pdl_old().get_field_idx<Tscal>(
"hpart"));
590 u32 len = vec_acc.size();
591 PatchDataField<Tscal> &f
592 = tmp.get_field<Tscal>(sched.pdl_old().get_field_idx<Tscal>(
"uint"));
593 sycl::buffer<Tscal> buf(vec_u.data(), len);
594 f.override(buf, len);
598 u32 len = vec_acc.size();
599 PatchDataField<Tvec> &f
600 = tmp.get_field<Tvec>(sched.pdl_old().get_field_idx<Tvec>(
"vxyz"));
601 sycl::buffer<Tvec> buf(vec_vel.data(), len);
602 f.override(buf, len);
605 pdat.insert_elements(tmp);
608 std::string log_gathered =
"";
612 logger::info_ln(
"Model",
"Push particles : ", log_gathered);
615 modules::ComputeLoadBalanceValue<Tvec, SPHKernel>(
616 ctx, solver.solver_config, solver.storage)
617 .update_load_balancing();
622 auto [m, M] = sched.get_box_tranform<Tvec>();
624 SerialPatchTree<Tvec> sptree(
629 shamrock::ReattributeDataUtility reatrib(sched);
634 reatrib.reatribute_patch_objects(sptree,
"xyz");
637 sched.check_patchdata_locality_correctness();
642 sched.for_each_local_patchdata([&](
const Patch &p, PatchDataLayer &pdat) {
643 log += shambase::format(
644 "\n patch id={}, N={} particles", p.id_patch, pdat.get_obj_cnt());
651 logger::info_ln(
"Model",
"current particle counts : ", log_gathered);
654 void remap_positions(std::function<Tvec(Tvec)> map);
657 std::vector<Tvec> &part_pos_insert,
658 std::vector<Tscal> &part_hpart_insert,
659 std::vector<Tscal> &part_u_insert);
661 void push_particle_mhd(
662 std::vector<Tvec> &part_pos_insert,
663 std::vector<Tscal> &part_hpart_insert,
664 std::vector<Tscal> &part_u_insert,
665 std::vector<Tvec> &part_B_on_rho_insert,
666 std::vector<Tscal> &part_psi_on_ch_insert);
669 inline void set_value_in_a_box(
670 std::string field_name, T val, std::pair<Tvec, Tvec> box,
u32 ivar) {
674 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
675 PatchDataField<Tvec> &
xyz
676 = pdat.template get_field<Tvec>(sched.pdl_old().
get_field_idx<Tvec>(
"xyz"));
679 = pdat.template get_field<T>(sched.pdl_old().
get_field_idx<T>(field_name));
681 if (ivar >= f.get_nvar()) {
683 "You are trying to set value in a box for field ({}) with "
684 "ivar ({}) >= f.get_nvar ({})",
690 u32 nvar = f.get_nvar();
693 auto acc = f.get_buf().template mirror_to<sham::host>();
694 auto acc_xyz =
xyz.get_buf().template mirror_to<sham::host>();
696 for (
u32 i = 0; i < f.get_obj_cnt(); i++) {
699 if (BBAA::is_coord_in_range(r, std::get<0>(box), std::get<1>(box))) {
700 acc[i * nvar + ivar] = val;
708 inline void set_value_in_sphere(std::string field_name, T val, Tvec center, Tscal radius) {
712 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
713 PatchDataField<Tvec> &
xyz
714 = pdat.template get_field<Tvec>(sched.pdl_old().
get_field_idx<Tvec>(
"xyz"));
717 = pdat.template get_field<T>(sched.pdl_old().
get_field_idx<T>(field_name));
719 if (f.get_nvar() != 1) {
723 Tscal r2 = radius * radius;
725 auto acc = f.get_buf().template mirror_to<sham::host>();
726 auto acc_xyz =
xyz.get_buf().template mirror_to<sham::host>();
728 for (
u32 i = 0; i < f.get_obj_cnt(); i++) {
729 Tvec dr = acc_xyz[i] - center;
731 if (sycl::dot(dr, dr) < r2) {
740 inline void add_kernel_value(std::string field_name, T val, Tvec center, Tscal h_ker) {
744 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
745 PatchDataField<Tvec> &
xyz
746 = pdat.template get_field<Tvec>(sched.pdl_old().
get_field_idx<Tvec>(
"xyz"));
749 = pdat.template get_field<T>(sched.pdl_old().
get_field_idx<T>(field_name));
751 if (f.get_nvar() != 1) {
756 auto acc = f.get_buf().template mirror_to<sham::host>();
757 auto acc_xyz =
xyz.get_buf().template mirror_to<sham::host>();
759 for (
u32 i = 0; i < f.get_obj_cnt(); i++) {
760 Tvec dr = acc_xyz[i] - center;
762 Tscal r = sycl::length(dr);
764 acc[i] += val * Kernel::W_3d(r, h_ker);
771 inline T get_sum(std::string name) {
773 T sum = shambase::VectorProperties<T>::get_zero();
777 [&](
u64 patch_id, shamrock::patch::PatchDataLayer &pdat) {
778 PatchDataField<T> &
xyz
779 = pdat.template get_field<T>(sched.pdl_old().
get_field_idx<T>(name));
781 sum +=
xyz.compute_sum();
784 return shamalgs::collective::allreduce_sum(sum);
787 Tvec get_closest_part_to(Tvec pos);
789 inline void apply_momentum_offset(Tvec offset) {
798 sched.for_each_patchdata_nonempty(
799 [&](shamrock::patch::Patch p, shamrock::patch::PatchDataLayer &pdat) {
800 tot_mass += solver.solver_config.gpart_mass * pdat.get_obj_cnt();
803 tot_mass = shamalgs::collective::allreduce_sum(tot_mass);
807 auto &mass = get_sink_mass<Tvec>(sync);
809 for (
size_t i = 0; i < mass.size(); i++) {
815 Tvec offset_vel = (tot_mass > 0) ? (offset / tot_mass)
816 : shambase::VectorProperties<Tvec>::get_zero();
819 auto &vel = get_sink_vel<Tvec>(sync);
821 for (
size_t i = 0; i < vel.size(); i++) {
822 vel[i] += offset_vel;
827 sched.for_each_patchdata_nonempty(
828 [&](shamrock::patch::Patch p, shamrock::patch::PatchDataLayer &pdat) {
829 PatchDataField<Tvec> &
vxyz = pdat.get_field<Tvec>(ivxyz);
830 vxyz.apply_offset(offset_vel);
834 inline void apply_position_offset(Tvec offset) {
843 for (
size_t i = 0; i < pos.size(); i++) {
849 sched.for_each_patchdata_nonempty(
850 [&](shamrock::patch::Patch p, shamrock::patch::PatchDataLayer &pdat) {
851 PatchDataField<Tvec> &
xyz = pdat.get_field<Tvec>(ixyz);
852 xyz.apply_offset(offset);
864 inline void set_solver_config(
typename Solver::Config cfg) {
865 if (ctx.is_scheduler_initialized()) {
867 "Cannot change solver config after scheduler is initialized");
870 solver.solver_config = cfg;
873 inline f64 solver_logs_last_rate() {
return solver.solve_logs.get_last_rate(); }
874 inline u64 solver_logs_last_obj_count() {
return solver.solve_logs.get_last_obj_count(); }
875 inline f64 solver_logs_cumulated_step_time() {
876 return solver.solve_logs.get_cumulated_step_time();
878 inline void solver_logs_reset_cumulated_step_time() {
879 solver.solve_logs.reset_cumulated_step_time();
881 inline u64 solver_logs_step_count() {
return solver.solve_logs.get_step_count(); }
882 inline void solver_logs_reset_step_count() { solver.solve_logs.reset_step_count(); }
884 inline void change_htolerances(Tscal in_coarse, Tscal in_fine) {
885 if (in_coarse < in_fine) {
887 "in_coarse ({}) must be greater than in_fine ({})", in_coarse, in_fine));
889 solver.solver_config.htol_up_coarse_cycle = in_coarse;
890 solver.solver_config.htol_up_fine_cycle = in_fine;
912 std::string metadata_user{};
916 nlohmann::json j = nlohmann::json::parse(metadata_user);
918 j.at(
"solver_config").get_to(solver.solver_config);
923 if (j.contains(
"sinks") && !j.at(
"sinks").is_null()) {
924 std::vector<SinkParticle<Tvec>> out;
925 j.at(
"sinks").get_to(out);
936 = std::find(sync_names.begin(), sync_names.end(),
"time") != sync_names.end();
939 solver.ensure_time_state_edges();
941 if (!had_time_edge) {
942 if (j.at(
"solver_config").contains(
"time_state")) {
946 "Migrated time/dt/cfl from solver_config.time_state into scheduler "
948 const auto &ts = j.at(
"solver_config").at(
"time_state");
949 solver.set_time(ts.at(
"time").get<Tscal>());
950 solver.set_next_dt(ts.at(
"dt_sph").get<Tscal>());
951 solver.set_cfl_multipler(ts.at(
"cfl_multiplier").get<Tscal>());
954 "this should never happen: dump has neither time edges nor "
955 "solver_config.time_state");
959 solver.init_ghost_layout();
961 solver.init_solver_graph();
963 shamlog_debug_ln(
"Sys",
"build local scheduler tables");
977 inline void dump(std::string fname) {
982 solver.update_sync_load_values();
984 nlohmann::json metadata;
985 metadata[
"solver_config"] = solver.solver_config;
997 f64 evolve_once_time_expl(
f64 t_curr,
f64 dt_input);
1001 inline void evolve_once() {
1002 solver.evolve_once();
1003 solver.print_timestep_logs();
1006 inline EvolveUntilResults evolve_until(
1007 Tscal target_time,
i32 niter_max,
f64 max_walltime = -1) {
1008 return solver.evolve_until(target_time, niter_max, max_walltime);
1012 void add_pdat_to_phantom_block(
1015 template<
class Tscal>
1016 inline void warp_disc(
1017 std::vector<Tvec> &pos,
1018 std::vector<Tvec> &vel,
1023 Tvec k = Tvec(-std::sin(posangle), std::cos(posangle), 0.);
1026 u32 len = pos.size();
1029 Tscal incl_rad = incl * shambase::constants::pi<Tscal> / 180.;
1031 for (
i32 i = 0; i < len; i++) {
1032 Tvec R_vec = pos[i];
1033 Tscal R = sycl::sqrt(sycl::dot(R_vec, R_vec));
1034 if (R < Rwarp - Hwarp) {
1036 }
else if (R < Rwarp + 3. * Hwarp && R > Rwarp - Hwarp) {
1040 + sycl::sin(shambase::constants::pi<Tscal> / (2. * Hwarp) * (R - Rwarp)))
1041 * sycl::sin(incl_rad));
1042 psi = shambase::constants::pi<Tscal>
1043 * Rwarp / (4. * Hwarp) * sycl::sin(incl_rad)
1044 / sycl::sqrt(1. - (0.5 * sycl::pow(sycl::sin(incl_rad), 2)));
1045 Tscal psimax = sycl::max(psimax, psi);
1046 Tscal x = pos[i].x();
1047 Tscal y = pos[i].y();
1048 Tscal z = pos[i].z();
1054 Tvec kk = Tvec(0., 0., 1.);
1055 Tvec w = sycl::cross(kk, pos[i]);
1057 pos[i] = pos[i] * sycl::cos(inc) + w * sycl::sin(inc)
1058 + kk * sycl::dot(kk, pos[i]) * (1. - sycl::cos(inc));
1066 inline void rotate_vector(Tvec &u, Tvec &v, Tscal theta) {
1068 Tvec vunit = v / sycl::sqrt(sycl::dot(v, v));
1069 Tvec w = sycl::cross(vunit, u);
1071 u = u * sycl::cos(theta) + w * sycl::sin(theta)
1072 + vunit * sycl::dot(vunit, u) * (1. - sycl::cos(theta));