Shamrock 2025.10.0
Astrophysical Code
Loading...
Searching...
No Matches
PatchDataToPy.hpp
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
10#pragma once
11
17
21#include <pybind11/numpy.h>
22#include <pybind11/pybind11.h>
23#include <optional>
24
25namespace shamrock {
26 template<class T>
27 class VecToNumpy;
28
29 template<class T>
30 class VecToNumpy {
31 public:
32 static py::array_t<T> convert(std::vector<T> vec) {
33
34 u32 len = vec.size();
35
36 py::array_t<T> ret({len});
37 auto r = ret.mutable_unchecked();
38
39 for (u32 i = 0; i < len; i++) {
40 r(i) = vec[i];
41 }
42
43 return std::move(ret);
44 }
45 };
46
47 template<class T>
48 class VecToNumpy<sycl::vec<T, 2>> {
49 public:
50 static py::array_t<T> convert(std::vector<sycl::vec<T, 2>> vec) {
51
52 u32 len = vec.size();
53
54 py::array_t<T> ret({len, 2U});
55 auto r = ret.mutable_unchecked();
56
57 for (u32 i = 0; i < len; i++) {
58 r(i, 0) = vec[i].x();
59 r(i, 1) = vec[i].y();
60 }
61 return std::move(ret);
62 }
63 };
64
65 template<class T>
66 class VecToNumpy<sycl::vec<T, 3>> {
67 public:
68 static py::array_t<T> convert(std::vector<sycl::vec<T, 3>> vec) {
69
70 u32 len = vec.size();
71
72 py::array_t<T> ret({len, 3U});
73 auto r = ret.mutable_unchecked();
74
75 for (u32 i = 0; i < len; i++) {
76 r(i, 0) = vec[i].x();
77 r(i, 1) = vec[i].y();
78 r(i, 2) = vec[i].z();
79 }
80 return std::move(ret);
81 }
82 };
83
84 template<class T>
85 class VecToNumpy<sycl::vec<T, 4>> {
86 public:
87 static py::array_t<T> convert(std::vector<sycl::vec<T, 4>> vec) {
88
89 u32 len = vec.size();
90
91 py::array_t<T> ret({len, 4U});
92 auto r = ret.mutable_unchecked();
93
94 for (u32 i = 0; i < len; i++) {
95 r(i, 0) = vec[i].x();
96 r(i, 1) = vec[i].y();
97 r(i, 2) = vec[i].z();
98 r(i, 3) = vec[i].w();
99 }
100 return std::move(ret);
101 }
102 };
103
104 template<class T>
105 class VecToNumpy<sycl::vec<T, 8>> {
106 public:
107 static py::array_t<T> convert(std::vector<sycl::vec<T, 8>> vec) {
108
109 u32 len = vec.size();
110
111 py::array_t<T> ret({len, 8U});
112 auto r = ret.mutable_unchecked();
113
114 for (u32 i = 0; i < len; i++) {
115 r(i, 0) = vec[i].s0();
116 r(i, 1) = vec[i].s1();
117 r(i, 2) = vec[i].s2();
118 r(i, 3) = vec[i].s3();
119 r(i, 4) = vec[i].s4();
120 r(i, 5) = vec[i].s5();
121 r(i, 6) = vec[i].s6();
122 r(i, 7) = vec[i].s7();
123 }
124 return std::move(ret);
125 }
126 };
127
128 template<class T>
129 class VecToNumpy<sycl::vec<T, 16>> {
130 public:
131 static py::array_t<T> convert(std::vector<sycl::vec<T, 16>> vec) {
132
133 u32 len = vec.size();
134
135 py::array_t<T> ret({len, 16U});
136 auto r = ret.mutable_unchecked();
137
138 for (u32 i = 0; i < len; i++) {
139 r(i, 0) = vec[i].s0();
140 r(i, 1) = vec[i].s1();
141 r(i, 2) = vec[i].s2();
142 r(i, 3) = vec[i].s3();
143 r(i, 4) = vec[i].s4();
144 r(i, 5) = vec[i].s5();
145 r(i, 6) = vec[i].s6();
146 r(i, 7) = vec[i].s7();
147 r(i, 8) = vec[i].s8();
148 r(i, 9) = vec[i].s9();
149 r(i, 10) = vec[i].sA();
150 r(i, 11) = vec[i].sB();
151 r(i, 12) = vec[i].sC();
152 r(i, 13) = vec[i].sD();
153 r(i, 14) = vec[i].sE();
154 r(i, 15) = vec[i].sF();
155 }
156
157 return std::move(ret);
158 }
159 };
160
161 template<class T>
162 void append_to_map(
163 std::string key,
164 std::vector<std::reference_wrapper<shamrock::patch::PatchDataLayer>> ref_lst,
165 py::dict &dic_out) {
166
167 std::vector<T> vec;
168
169 auto appender = [&](auto &field) {
170 if (field.get_name() == key) {
171
172 logger::debug_ln("PatchDataToPy", "appending field", key);
173
174 {
175 auto acc = field.get_buf().copy_to_stdvec();
176 u32 len = field.get_val_cnt();
177
178 for (u32 i = 0; i < len; i++) {
179 vec.push_back(acc[i]);
180 }
181 }
182 }
183 };
184
185 for (auto &pdat_ref : ref_lst) {
186 auto &pdat = pdat_ref.get();
187 if (pdat.get_obj_cnt() > 0) {
188 pdat.for_each_field<T>([&](auto &field) {
189 appender(field);
190 });
191 }
192 }
193
194 if (!vec.empty()) {
195 auto arr = VecToNumpy<T>::convert(vec);
196
197 logger::debug_ln("PatchDataToPy", "adding -> ", key);
198
199 if (dic_out.contains(key.c_str())) {
200 throw shambase::make_except_with_loc<std::runtime_error>("the key already exists");
201 } else {
202 dic_out[key.c_str()] = arr;
203 }
204 }
205 }
206
207 template<class T>
208 void append_to_map(
209 std::string key,
210 std::vector<std::unique_ptr<shamrock::patch::PatchDataLayer>> &lst,
211 py::dict &dic_out) {
212
213 std::vector<std::reference_wrapper<shamrock::patch::PatchDataLayer>> ref_lst;
214 for (auto &pdat : lst) {
215 if (pdat) {
216 ref_lst.push_back(*pdat);
217 }
218 }
219
220 append_to_map<T>(key, ref_lst, dic_out);
221 }
222
223 inline py::dict pdat_to_dic(shamrock::patch::PatchDataLayer &pdat) {
224 py::dict dic_out;
225
226 std::reference_wrapper<shamrock::patch::PatchDataLayer> ref_pdat = pdat;
227
228 using namespace shamrock;
229
230 for (auto fname : pdat.pdl().get_field_names()) {
231 append_to_map<f32>(fname, {ref_pdat}, dic_out);
232 append_to_map<f32_2>(fname, {ref_pdat}, dic_out);
233 append_to_map<f32_3>(fname, {ref_pdat}, dic_out);
234 append_to_map<f32_4>(fname, {ref_pdat}, dic_out);
235 append_to_map<f32_8>(fname, {ref_pdat}, dic_out);
236 append_to_map<f32_16>(fname, {ref_pdat}, dic_out);
237 append_to_map<f64>(fname, {ref_pdat}, dic_out);
238 append_to_map<f64_2>(fname, {ref_pdat}, dic_out);
239 append_to_map<f64_3>(fname, {ref_pdat}, dic_out);
240 append_to_map<f64_4>(fname, {ref_pdat}, dic_out);
241 append_to_map<f64_8>(fname, {ref_pdat}, dic_out);
242 append_to_map<f64_16>(fname, {ref_pdat}, dic_out);
243 append_to_map<u32>(fname, {ref_pdat}, dic_out);
244 append_to_map<u64>(fname, {ref_pdat}, dic_out);
245 append_to_map<u32_3>(fname, {ref_pdat}, dic_out);
246 append_to_map<u64_3>(fname, {ref_pdat}, dic_out);
247 append_to_map<i64_3>(fname, {ref_pdat}, dic_out);
248 }
249
250 return dic_out;
251 }
252
253 namespace details {
254 template<class T>
255 inline bool try_get_field_as_np(
256 shamrock::patch::PatchDataLayer &pdat,
257 const std::string &key,
258 std::optional<py::object> &ret) {
259
260 pdat.for_each_field<T>([&](auto &field) {
261 if (ret.has_value() || field.get_name() != key) {
262 return;
263 }
264 ret = VecToNumpy<T>::convert(field.get_buf().copy_to_stdvec());
265 });
266
267 return ret.has_value();
268 }
269 } // namespace details
270
279 class PatchDataLazyGetter {
280 public:
282
283 explicit PatchDataLazyGetter(shamrock::patch::PatchDataLayer &pdat) : pdat(pdat) {}
284
285 py::object get_item(const std::string &key) const {
286 std::optional<py::object> ret;
287
288 // clang-format off
289 details::try_get_field_as_np<f32>(pdat, key, ret)
290 || details::try_get_field_as_np<f32_2>(pdat, key, ret)
291 || details::try_get_field_as_np<f32_3>(pdat, key, ret)
292 || details::try_get_field_as_np<f32_4>(pdat, key, ret)
293 || details::try_get_field_as_np<f32_8>(pdat, key, ret)
294 || details::try_get_field_as_np<f32_16>(pdat, key, ret)
295 || details::try_get_field_as_np<f64>(pdat, key, ret)
296 || details::try_get_field_as_np<f64_2>(pdat, key, ret)
297 || details::try_get_field_as_np<f64_3>(pdat, key, ret)
298 || details::try_get_field_as_np<f64_4>(pdat, key, ret)
299 || details::try_get_field_as_np<f64_8>(pdat, key, ret)
300 || details::try_get_field_as_np<f64_16>(pdat, key, ret)
301 || details::try_get_field_as_np<u32>(pdat, key, ret)
302 || details::try_get_field_as_np<u64>(pdat, key, ret)
303 || details::try_get_field_as_np<u32_3>(pdat, key, ret)
304 || details::try_get_field_as_np<u64_3>(pdat, key, ret)
305 || details::try_get_field_as_np<i64_3>(pdat, key, ret);
306 // clang-format on
307
308 if (!ret.has_value()) {
309 throw py::key_error(key);
310 }
311
312 return *ret;
313 }
314 };
315} // namespace shamrock
std::uint32_t u32
32 bit unsigned integer
std::vector< std::string > get_field_names()
Get the list of field names.
PatchDataLayer container class, the layout is described in patchdata_layout.
ExcptTypes make_except_with_loc(std::string message, SourceLocation loc=SourceLocation{})
Create an exception with a message and a location.
namespace for the main framework
Definition __init__.py:1
Pybind11 include and definitions.
void debug_ln(std::string module_name, Types... var2)
Prints a log message with multiple arguments followed by a newline.
Definition logs.hpp:132