43 auto edges = get_edges();
45 edges.spans_positions.check_sizes(edges.sizes.indexes);
46 edges.spans_velocities.check_sizes(edges.sizes.indexes);
47 edges.spans_accelerations.check_sizes(edges.sizes.indexes);
48 edges.spans_masses.check_sizes(edges.sizes.indexes);
49 edges.spans_accel_ext.check_sizes(edges.sizes.indexes);
51 const Tvec x0 = edges.central_pos.data;
52 const Tvec v0 = edges.central_vel.data;
53 const Tvec a0 = edges.central_acc.data;
54 const Tscal fac = edges.gw_prefactor.data;
55 const Tscal theta_deg = edges.theta_gw.data;
56 const Tscal phi_deg = edges.phi_gw.data;
58 constexpr Tscal pi = shambase::constants::pi<Tscal>;
60 auto dev_sched = shamsys::instance::get_compute_scheduler_ptr();
66 auto make_ddq = [&]() {
67 return edges.sizes.indexes.template map<sham::DeviceBuffer<Tscal>>(
68 [&](
u64 id,
const u32 &n) {
83 edges.spans_positions.get_spans(),
84 edges.spans_velocities.get_spans(),
85 edges.spans_accelerations.get_spans(),
86 edges.spans_masses.get_spans(),
87 edges.spans_accel_ext.get_spans()},
103 const Tscal m = mass[gid];
105 const Tscal x = xyz[gid][0] - x0[0];
106 const Tscal y = xyz[gid][1] - x0[1];
107 const Tscal z = xyz[gid][2] - x0[2];
108 const Tscal vx = vxyz[gid][0] - v0[0];
109 const Tscal vy = vxyz[gid][1] - v0[1];
110 const Tscal vz = vxyz[gid][2] - v0[2];
113 const Tscal ax = axyz[gid][0] - a0[0];
114 const Tscal ay = axyz[gid][1] - a0[1];
115 const Tscal az = axyz[gid][2] - a0[2];
117 ddq0[gid] = m * (Tscal(2.) * vx * vx + x * ax + x * ax);
118 ddq1[gid] = m * (Tscal(2.) * vx * vy + x * ay + y * ax);
119 ddq2[gid] = m * (Tscal(2.) * vx * vz + x * az + z * ax);
120 ddq3[gid] = m * (Tscal(2.) * vy * vy + y * ay + y * ay);
121 ddq4[gid] = m * (Tscal(2.) * vy * vz + y * az + z * ay);
122 ddq5[gid] = m * (Tscal(2.) * vz * vz + z * az + z * az);
127 Tscal ddq0_per_rank = 0, ddq1_per_rank = 0, ddq2_per_rank = 0, ddq3_per_rank = 0,
128 ddq4_per_rank = 0, ddq5_per_rank = 0;
130 edges.sizes.indexes.for_each([&](
u64 id,
const u32 &n) {
139 edges.ddq.data[0] = shamalgs::collective::allreduce_sum(ddq0_per_rank);
140 edges.ddq.data[1] = shamalgs::collective::allreduce_sum(ddq1_per_rank);
141 edges.ddq.data[2] = shamalgs::collective::allreduce_sum(ddq2_per_rank);
142 edges.ddq.data[3] = shamalgs::collective::allreduce_sum(ddq3_per_rank);
143 edges.ddq.data[4] = shamalgs::collective::allreduce_sum(ddq4_per_rank);
144 edges.ddq.data[5] = shamalgs::collective::allreduce_sum(ddq5_per_rank);
146 std::array<Tscal, 9> Q_arr{};
147 Mat3<Tscal> Q(Q_arr.data());
148 Q(0, 0) = edges.ddq.data[0];
149 Q(0, 1) = Q(1, 0) = edges.ddq.data[1];
150 Q(0, 2) = Q(2, 0) = edges.ddq.data[2];
151 Q(1, 1) = edges.ddq.data[3];
152 Q(1, 2) = Q(2, 1) = edges.ddq.data[4];
153 Q(2, 2) = edges.ddq.data[5];
155 std::array<Tscal, 9> ddq_xy_arr{};
156 Mat3<Tscal> ddq_xy(ddq_xy_arr.data());
158 const bool rotate = std::abs(theta_deg) >
static_cast<Tscal
>(1e-30);
160 const Tscal lam = theta_deg * pi /
static_cast<Tscal
>(180);
161 const Tscal c = std::cos(lam);
162 const Tscal s = std::sin(lam);
164 std::array<Tscal, 9> R_arr
165 = {c, Tscal(0), s, Tscal(0), Tscal(1), Tscal(0), -s, Tscal(0), c};
166 Mat3<Tscal> R(R_arr.data());
168 std::array<Tscal, 9> inter_arr{};
169 Mat3<Tscal> inter(inter_arr.data());
170 for (std::size_t i = 0; i < 3; ++i) {
171 for (std::size_t j = 0; j < 3; ++j) {
172 Tscal sum = Tscal(0);
173 for (std::size_t k = 0; k < 3; ++k) {
174 sum += Q(i, k) * R(k, j);
180 for (std::size_t i = 0; i < 3; ++i) {
181 for (std::size_t j = 0; j < 3; ++j) {
182 Tscal sum = Tscal(0);
183 for (std::size_t k = 0; k < 3; ++k) {
184 sum += R(k, i) * inter(k, j);
190 for (std::size_t i = 0; i < 9; ++i) {
191 ddq_xy_arr[i] = Q_arr[i];
196 const Tscal phi = phi_deg * pi /
static_cast<Tscal
>(180);
197 const Tscal sinphi = std::sin(phi);
198 const Tscal cosphi = std::cos(phi);
199 const Tscal sinphi2 = sinphi * sinphi;
200 const Tscal cosphi2 = cosphi * cosphi;
201 const Tscal sin2phi = std::sin(Tscal(2) * phi);
202 const Tscal cos2phi = std::cos(Tscal(2) * phi);
204 std::array<Tscal, 4> hx_out{};
205 std::array<Tscal, 4> hp_out{};
207 for (
u32 i = 0; i < 4; ++i) {
208 const Tscal eta =
static_cast<Tscal
>(i) * pi /
static_cast<Tscal
>(6);
209 const Tscal sineta = std::sin(eta);
210 const Tscal coseta = std::cos(eta);
211 const Tscal sineta2 = sineta * sineta;
212 const Tscal coseta2 = coseta * coseta;
213 const Tscal sin2eta = std::sin(Tscal(2) * eta);
216 * (ddq_xy(0, 0) * (cosphi2 - sinphi2 * coseta2)
217 + ddq_xy(1, 1) * (sinphi2 - cosphi2 * coseta2) - ddq_xy(2, 2) * sineta2
218 - ddq_xy(0, 1) * sin2phi * (Tscal(1) + coseta2)
219 + ddq_xy(0, 2) * sinphi * sin2eta + ddq_xy(1, 2) * cosphi * sin2eta);
221 hx_out[i] = Tscal(2) * fac
222 * (Tscal(0.5) * (ddq_xy(0, 0) - ddq_xy(1, 1)) * sin2phi * coseta
223 + ddq_xy(0, 1) * cos2phi * coseta - ddq_xy(0, 2) * cosphi * sineta
224 + ddq_xy(1, 2) * sinphi * sineta);
227 edges.ddq_xy.data = ddq_xy_arr;
228 edges.hx.data = hx_out;
229 edges.hp.data = hp_out;