DDC 0.16.0
Loading...
Searching...
No Matches
splines_linear_problem_band.cpp
1// Copyright (C) The DDC development team, see COPYRIGHT.md file
2//
3// SPDX-License-Identifier: MIT
4
5#include <algorithm>
6#include <cassert>
7#if !defined(NDEBUG)
8# include <cmath>
9# include <limits>
10#endif
11#include <cstddef>
12#include <stdexcept>
13#include <string>
14#include <type_traits>
15
16#include <ddc/ddc.hpp>
17
18#include <Kokkos_Core.hpp>
19
20#if __has_include(<mkl_lapacke.h>)
21# include <mkl_lapacke.h>
22#else
23# include <lapacke.h>
24#endif
25
26#include <KokkosBatched_Gbtrs.hpp>
27#include <KokkosBatched_Util.hpp>
28
31
32namespace ddc::detail {
33
34template <class ExecSpace>
35SplinesLinearProblemBand<ExecSpace>::SplinesLinearProblemBand(
36 std::size_t const mat_size,
37 std::size_t const kl,
38 std::size_t const ku)
39 : SplinesLinearProblem<ExecSpace>(mat_size)
40 , m_kl(kl)
41 , m_ku(ku)
42 /*
43 * The matrix itself stored in band format requires a (kl + ku + 1)*mat_size
44 * allocation, but the LU-factorization requires an additional kl*mat_size block
45 */
46 , m_q("q", 2 * kl + ku + 1, mat_size)
47 , m_ipiv("ipiv", mat_size)
48{
49 assert(m_kl <= mat_size);
50 assert(m_ku <= mat_size);
51
52 Kokkos::deep_copy(m_q.view_host(), 0.);
53}
54
55template <class ExecSpace>
56SplinesLinearProblemBand<ExecSpace>::~SplinesLinearProblemBand() = default;
57
58template <class ExecSpace>
59std::size_t SplinesLinearProblemBand<ExecSpace>::band_storage_row_index(
60 std::size_t const i,
61 std::size_t const j) const
62{
63 return m_kl + m_ku + i - j;
64}
65
66template <class ExecSpace>
67Real SplinesLinearProblemBand<ExecSpace>::get_element(std::size_t const i, std::size_t const j)
68 const
69{
70 assert(i < size());
71 assert(j < size());
72 /*
73 * The "row index" of the band format storage identify the (sub/super)-diagonal
74 * while the column index is actually the column index of the matrix. Two layouts
75 * are supported by LAPACKE. The m_kl first lines are irrelevant for the storage of
76 * the matrix itself but required for the storage of its LU factorization.
77 */
78 if (i >= std::
79 max(static_cast<std::ptrdiff_t>(0),
80 static_cast<std::ptrdiff_t>(j) - static_cast<std::ptrdiff_t>(m_ku))
81 && i < std::min(size(), j + m_kl + 1)) {
82 return m_q.view_host()(band_storage_row_index(i, j), j);
83 }
84
85 return 0.0;
86}
87
88template <class ExecSpace>
89void SplinesLinearProblemBand<ExecSpace>::set_element(
90 std::size_t const i,
91 std::size_t const j,
92 Real const aij)
93{
94 assert(i < size());
95 assert(j < size());
96 /*
97 * The "row index" of the band format storage identify the (sub/super)-diagonal
98 * while the column index is actually the column index of the matrix. Two layouts
99 * are supported by LAPACKE. The m_kl first lines are irrelevant for the storage of
100 * the matrix itself but required for the storage of its LU factorization.
101 */
102 if (i >= std::
103 max(static_cast<std::ptrdiff_t>(0),
104 static_cast<std::ptrdiff_t>(j) - static_cast<std::ptrdiff_t>(m_ku))
105 && i < std::min(size(), j + m_kl + 1)) {
106 m_q.view_host()(band_storage_row_index(i, j), j) = aij;
107 } else {
108 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
109 }
110}
111
112template <class ExecSpace>
113void SplinesLinearProblemBand<ExecSpace>::setup_solver()
114{
115 int info;
116 if constexpr (std::is_same_v<Real, float>) {
117 info = LAPACKE_sgbtrf(
118 LAPACK_ROW_MAJOR,
119 size(),
120 size(),
121 m_kl,
122 m_ku,
123 m_q.view_host().data(),
124 m_q.view_host().stride(
125 0), // m_q.view_host().stride(0) if LAPACK_ROW_MAJOR, m_q.view_host().stride(1) if LAPACK_COL_MAJOR
126 m_ipiv.view_host().data());
127 } else {
128 info = LAPACKE_dgbtrf(
129 LAPACK_ROW_MAJOR,
130 size(),
131 size(),
132 m_kl,
133 m_ku,
134 m_q.view_host().data(),
135 m_q.view_host().stride(
136 0), // m_q.view_host().stride(0) if LAPACK_ROW_MAJOR, m_q.view_host().stride(1) if LAPACK_COL_MAJOR
137 m_ipiv.view_host().data());
138 }
139 if (info != 0) {
140 throw std::runtime_error("LAPACKE_gbtrf failed with error code " + std::to_string(info));
141 }
142
143 // Convert 1-based index to 0-based index
144 for (std::size_t i = 0; i < size(); ++i) {
145 m_ipiv.view_host()(i) -= 1;
146 }
147
148 // Push on device
149 m_q.modify_host();
150 m_q.sync_device();
151 m_ipiv.modify_host();
152 m_ipiv.sync_device();
153}
154
155template <class ExecSpace>
156void SplinesLinearProblemBand<ExecSpace>::solve(MultiRHS const b, bool const transpose) const
157{
158 assert(b.extent(0) == size());
159
160 std::size_t const kl_proxy = m_kl;
161 std::size_t const ku_proxy = m_ku;
162 auto q_device = m_q.view_device();
163 auto ipiv_device = m_ipiv.view_device();
164 Kokkos::RangePolicy<ExecSpace> const policy(0, b.extent(1));
165 if (transpose) {
166 Kokkos::parallel_for(
167 "gbtrs",
168 policy,
169 KOKKOS_LAMBDA(int const i) {
170 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
171 KokkosBatched::SerialGbtrs<
172 KokkosBatched::Trans::Transpose,
173 KokkosBatched::Algo::Gbtrs::Unblocked>::
174 invoke(q_device, ipiv_device, sub_b, kl_proxy, ku_proxy);
175 });
176 } else {
177 Kokkos::parallel_for(
178 "gbtrs",
179 policy,
180 KOKKOS_LAMBDA(int const i) {
181 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
182 KokkosBatched::SerialGbtrs<
183 KokkosBatched::Trans::NoTranspose,
184 KokkosBatched::Algo::Gbtrs::Unblocked>::
185 invoke(q_device, ipiv_device, sub_b, kl_proxy, ku_proxy);
186 });
187 }
188}
189
190#if defined(KOKKOS_ENABLE_SERIAL)
191template class SplinesLinearProblemBand<Kokkos::Serial>;
192#endif
193#if defined(KOKKOS_ENABLE_OPENMP)
194template class SplinesLinearProblemBand<Kokkos::OpenMP>;
195#endif
196#if defined(KOKKOS_ENABLE_CUDA)
197template class SplinesLinearProblemBand<Kokkos::Cuda>;
198#endif
199#if defined(KOKKOS_ENABLE_HIP)
200template class SplinesLinearProblemBand<Kokkos::HIP>;
201#endif
202#if defined(KOKKOS_ENABLE_SYCL)
203template class SplinesLinearProblemBand<Kokkos::SYCL>;
204#endif
205
206} // namespace ddc::detail
The top-level namespace of DDC.