Shamrock 2025.10.0
Astrophysical Code
Loading...
Searching...
No Matches
SyclMpiTypes.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
15
18#include "shamcomm/logs.hpp"
20
21bool __mpi_sycl_type_active = false;
22bool is_mpi_sycl_interop_active() { return __mpi_sycl_type_active; }
23
24/*
25const int __len_vec2 [] = {1,1};
26const int __len_vec3 [] = {1,1,1};
27const int __len_vec4 [] = {1,1,1,1};
28const int __len_vec8 [] = {1,1,1,1,1,1,1,1};
29const int __len_vec16 [] = {1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1};
30*/
31const int __len_vec2 = 2;
32const int __len_vec3 = 3;
33const int __len_vec4 = 4;
34const int __len_vec8 = 8;
35const int __len_vec16 = 16;
36
37inline MPI_Datatype __tmp_mpi_type_i64_3;
38inline MPI_Datatype __tmp_mpi_type_i32_3;
39inline MPI_Datatype __tmp_mpi_type_i16_3;
40inline MPI_Datatype __tmp_mpi_type_i8_3;
41inline MPI_Datatype __tmp_mpi_type_u64_3;
42inline MPI_Datatype __tmp_mpi_type_u32_3;
43inline MPI_Datatype __tmp_mpi_type_u16_3;
44inline MPI_Datatype __tmp_mpi_type_u8_3;
45inline MPI_Datatype __tmp_mpi_type_f16_3;
46inline MPI_Datatype __tmp_mpi_type_f32_3;
47inline MPI_Datatype __tmp_mpi_type_f64_3;
48
49#define __SYCL_TYPE_COMMIT_len2(base_name, src_type) \
50 { \
51 check_offset_validity<base_name>(); \
52 MPICHECK(MPI_Type_contiguous(__len_vec2, mpi_type_##src_type, &mpi_type_##base_name)); \
53 MPICHECK(MPI_Type_commit(&mpi_type_##base_name)); \
54 shamlog_debug_mpi_ln("SyclMpiTypes", "init mpi type for : " #base_name); \
55 }
56
57#define __SYCL_TYPE_COMMIT_len3(base_name, src_type) \
58 { \
59 check_offset_validity<base_name>(); \
60 MPICHECK( \
61 MPI_Type_contiguous(__len_vec3, mpi_type_##src_type, &__tmp_mpi_type_##base_name)); \
62 MPICHECK(MPI_Type_create_resized( \
63 __tmp_mpi_type_##base_name, 0, sizeof(base_name), &mpi_type_##base_name)); \
64 MPICHECK(MPI_Type_commit(&mpi_type_##base_name)); \
65 shamlog_debug_mpi_ln("SyclMpiTypes", "init mpi type for : " #base_name); \
66 }
67
68#define __SYCL_TYPE_COMMIT_len4(base_name, src_type) \
69 { \
70 check_offset_validity<base_name>(); \
71 MPICHECK(MPI_Type_contiguous(__len_vec4, mpi_type_##src_type, &mpi_type_##base_name)); \
72 MPICHECK(MPI_Type_commit(&mpi_type_##base_name)); \
73 shamlog_debug_mpi_ln("SyclMpiTypes", "init mpi type for : " #base_name); \
74 }
75
76#define __SYCL_TYPE_COMMIT_len8(base_name, src_type) \
77 { \
78 check_offset_validity<base_name>(); \
79 MPICHECK(MPI_Type_contiguous(__len_vec8, mpi_type_##src_type, &mpi_type_##base_name)); \
80 MPICHECK(MPI_Type_commit(&mpi_type_##base_name)); \
81 shamlog_debug_mpi_ln("SyclMpiTypes", "init mpi type for : " #base_name); \
82 }
83
84#define __SYCL_TYPE_COMMIT_len16(base_name, src_type) \
85 { \
86 check_offset_validity<base_name>(); \
87 MPICHECK(MPI_Type_contiguous(__len_vec16, mpi_type_##src_type, &mpi_type_##base_name)); \
88 MPICHECK(MPI_Type_commit(&mpi_type_##base_name)); \
89 shamlog_debug_mpi_ln("SyclMpiTypes", "init mpi type for : " #base_name); \
90 }
91
92template<class T>
93void check_offset_validity() {
94 T a{};
95
96 std::ptrdiff_t base = reinterpret_cast<std::ptrdiff_t>(&a);
97 std::ptrdiff_t s0 = reinterpret_cast<std::ptrdiff_t>(&a.s0());
98
99 if (s0 - base != 0) {
101 "Offset is not valid for type {}, base = {}, s0 = {}", typeid(T).name(), base, s0));
102 }
103}
104
105void create_sycl_mpi_types() {
106
107 __SYCL_TYPE_COMMIT_len2(i64_2, i64);
108 __SYCL_TYPE_COMMIT_len2(i32_2, i32);
109 __SYCL_TYPE_COMMIT_len2(i16_2, i16);
110 __SYCL_TYPE_COMMIT_len2(i8_2, i8);
111 __SYCL_TYPE_COMMIT_len2(u64_2, u64);
112 __SYCL_TYPE_COMMIT_len2(u32_2, u32);
113 __SYCL_TYPE_COMMIT_len2(u16_2, u16);
114 __SYCL_TYPE_COMMIT_len2(u8_2, u8);
115 __SYCL_TYPE_COMMIT_len2(f16_2, f16);
116 __SYCL_TYPE_COMMIT_len2(f32_2, f32);
117 __SYCL_TYPE_COMMIT_len2(f64_2, f64);
118
119 __SYCL_TYPE_COMMIT_len3(i64_3, i64);
120 __SYCL_TYPE_COMMIT_len3(i32_3, i32);
121 __SYCL_TYPE_COMMIT_len3(i16_3, i16);
122
123 // {
124 // i16_3 a;
125
126 // MPI_Datatype types_list[3] = {mpi_type_i16,mpi_type_i16,mpi_type_i16};
127 // int block_lens[3] = {1,1,1};
128 // MPI_Aint MPI_offset[3];
129 // MPI_offset[0] = ((size_t) ( (char *)&(a.x()) - (char *)&(a) ));
130 // MPI_offset[1] = ((size_t) ( (char *)&(a.y()) - (char *)&(a) ));
131 // MPI_offset[2] = ((size_t) ( (char *)&(a.z()) - (char *)&(a) ));
132
133 // mpi::type_create_struct( 3, block_lens, MPI_offset, types_list, &mpi_type_i16_3 );
134 // /*mpi::type_create_resized(__tmp_mpi_type_i16_3, 0, sizeof(base_name),
135 // &mpi_type_i16_3);*/\ mpi::type_commit( &mpi_type_i16_3 );
136 // }
137
138 __SYCL_TYPE_COMMIT_len3(i8_3, i8);
139 __SYCL_TYPE_COMMIT_len3(u64_3, u64);
140 __SYCL_TYPE_COMMIT_len3(u32_3, u32);
141 __SYCL_TYPE_COMMIT_len3(u16_3, u16);
142 __SYCL_TYPE_COMMIT_len3(u8_3, u8);
143 __SYCL_TYPE_COMMIT_len3(f16_3, f16);
144 __SYCL_TYPE_COMMIT_len3(f32_3, f32);
145 __SYCL_TYPE_COMMIT_len3(f64_3, f64);
146
147 __SYCL_TYPE_COMMIT_len4(i64_4, i64);
148 __SYCL_TYPE_COMMIT_len4(i32_4, i32);
149 __SYCL_TYPE_COMMIT_len4(i16_4, i16);
150 __SYCL_TYPE_COMMIT_len4(i8_4, i8);
151 __SYCL_TYPE_COMMIT_len4(u64_4, u64);
152 __SYCL_TYPE_COMMIT_len4(u32_4, u32);
153 __SYCL_TYPE_COMMIT_len4(u16_4, u16);
154 __SYCL_TYPE_COMMIT_len4(u8_4, u8);
155 __SYCL_TYPE_COMMIT_len4(f16_4, f16);
156 __SYCL_TYPE_COMMIT_len4(f32_4, f32);
157 __SYCL_TYPE_COMMIT_len4(f64_4, f64);
158
159 __SYCL_TYPE_COMMIT_len8(i64_8, i64);
160 __SYCL_TYPE_COMMIT_len8(i32_8, i32);
161 __SYCL_TYPE_COMMIT_len8(i16_8, i16);
162 __SYCL_TYPE_COMMIT_len8(i8_8, i8);
163 __SYCL_TYPE_COMMIT_len8(u64_8, u64);
164 __SYCL_TYPE_COMMIT_len8(u32_8, u32);
165 __SYCL_TYPE_COMMIT_len8(u16_8, u16);
166 __SYCL_TYPE_COMMIT_len8(u8_8, u8);
167 __SYCL_TYPE_COMMIT_len8(f16_8, f16);
168 __SYCL_TYPE_COMMIT_len8(f32_8, f32);
169 __SYCL_TYPE_COMMIT_len8(f64_8, f64);
170
171 __SYCL_TYPE_COMMIT_len16(i64_16, i64);
172 __SYCL_TYPE_COMMIT_len16(i32_16, i32);
173 __SYCL_TYPE_COMMIT_len16(i16_16, i16);
174 __SYCL_TYPE_COMMIT_len16(i8_16, i8);
175 __SYCL_TYPE_COMMIT_len16(u64_16, u64);
176 __SYCL_TYPE_COMMIT_len16(u32_16, u32);
177 __SYCL_TYPE_COMMIT_len16(u16_16, u16);
178 __SYCL_TYPE_COMMIT_len16(u8_16, u8);
179 __SYCL_TYPE_COMMIT_len16(f16_16, f16);
180 __SYCL_TYPE_COMMIT_len16(f32_16, f32);
181 __SYCL_TYPE_COMMIT_len16(f64_16, f64);
182
183 __mpi_sycl_type_active = true;
184}
185
186void free_sycl_mpi_types() {
187
188 MPICHECK(MPI_Type_free(&mpi_type_i64_2));
189 MPICHECK(MPI_Type_free(&mpi_type_i32_2));
190 MPICHECK(MPI_Type_free(&mpi_type_i16_2));
191 MPICHECK(MPI_Type_free(&mpi_type_i8_2));
192 MPICHECK(MPI_Type_free(&mpi_type_u64_2));
193 MPICHECK(MPI_Type_free(&mpi_type_u32_2));
194 MPICHECK(MPI_Type_free(&mpi_type_u16_2));
195 MPICHECK(MPI_Type_free(&mpi_type_u8_2));
196 MPICHECK(MPI_Type_free(&mpi_type_f16_2));
197 MPICHECK(MPI_Type_free(&mpi_type_f32_2));
198 MPICHECK(MPI_Type_free(&mpi_type_f64_2));
199
200 MPICHECK(MPI_Type_free(&mpi_type_i64_3));
201 MPICHECK(MPI_Type_free(&mpi_type_i32_3));
202 MPICHECK(MPI_Type_free(&mpi_type_i16_3));
203 MPICHECK(MPI_Type_free(&mpi_type_i8_3));
204 MPICHECK(MPI_Type_free(&mpi_type_u64_3));
205 MPICHECK(MPI_Type_free(&mpi_type_u32_3));
206 MPICHECK(MPI_Type_free(&mpi_type_u16_3));
207 MPICHECK(MPI_Type_free(&mpi_type_u8_3));
208 MPICHECK(MPI_Type_free(&mpi_type_f16_3));
209 MPICHECK(MPI_Type_free(&mpi_type_f32_3));
210 MPICHECK(MPI_Type_free(&mpi_type_f64_3));
211
212 MPICHECK(MPI_Type_free(&mpi_type_i64_4));
213 MPICHECK(MPI_Type_free(&mpi_type_i32_4));
214 MPICHECK(MPI_Type_free(&mpi_type_i16_4));
215 MPICHECK(MPI_Type_free(&mpi_type_i8_4));
216 MPICHECK(MPI_Type_free(&mpi_type_u64_4));
217 MPICHECK(MPI_Type_free(&mpi_type_u32_4));
218 MPICHECK(MPI_Type_free(&mpi_type_u16_4));
219 MPICHECK(MPI_Type_free(&mpi_type_u8_4));
220 MPICHECK(MPI_Type_free(&mpi_type_f16_4));
221 MPICHECK(MPI_Type_free(&mpi_type_f32_4));
222 MPICHECK(MPI_Type_free(&mpi_type_f64_4));
223
224 MPICHECK(MPI_Type_free(&mpi_type_i64_8));
225 MPICHECK(MPI_Type_free(&mpi_type_i32_8));
226 MPICHECK(MPI_Type_free(&mpi_type_i16_8));
227 MPICHECK(MPI_Type_free(&mpi_type_i8_8));
228 MPICHECK(MPI_Type_free(&mpi_type_u64_8));
229 MPICHECK(MPI_Type_free(&mpi_type_u32_8));
230 MPICHECK(MPI_Type_free(&mpi_type_u16_8));
231 MPICHECK(MPI_Type_free(&mpi_type_u8_8));
232 MPICHECK(MPI_Type_free(&mpi_type_f16_8));
233 MPICHECK(MPI_Type_free(&mpi_type_f32_8));
234 MPICHECK(MPI_Type_free(&mpi_type_f64_8));
235
236 MPICHECK(MPI_Type_free(&mpi_type_i64_16));
237 MPICHECK(MPI_Type_free(&mpi_type_i32_16));
238 MPICHECK(MPI_Type_free(&mpi_type_i16_16));
239 MPICHECK(MPI_Type_free(&mpi_type_i8_16));
240 MPICHECK(MPI_Type_free(&mpi_type_u64_16));
241 MPICHECK(MPI_Type_free(&mpi_type_u32_16));
242 MPICHECK(MPI_Type_free(&mpi_type_u16_16));
243 MPICHECK(MPI_Type_free(&mpi_type_u8_16));
244 MPICHECK(MPI_Type_free(&mpi_type_f16_16));
245 MPICHECK(MPI_Type_free(&mpi_type_f32_16));
246 MPICHECK(MPI_Type_free(&mpi_type_f64_16));
247
248 __mpi_sycl_type_active = false;
249}
double f64
Alias for double.
float f32
Alias for float.
std::int8_t i8
8 bit integer
std::uint8_t u8
8 bit unsigned integer
std::uint32_t u32
32 bit unsigned integer
std::uint64_t u64
64 bit unsigned integer
std::uint16_t u16
16 bit unsigned integer
std::int16_t i16
16 bit integer
std::int64_t i64
64 bit integer
std::int32_t i32
32 bit integer
This header file contains utility functions related to exception handling in the code.
Utility functions for MPI error checking.
#define MPICHECK(mpicall)
Shortcut macro to check MPI return codes.
ExcptTypes make_except_with_loc(std::string message, SourceLocation loc=SourceLocation{})
Create an exception with a message and a location.