DDC 0.16.0
Loading...
Searching...
No Matches
splines_linear_problem_dense.cpp
1// Copyright (C) The DDC development team, see COPYRIGHT.md file
2//
3// SPDX-License-Identifier: MIT
4
5#include <cassert>
6#include <cstddef>
7#include <stdexcept>
8#include <string>
9#include <type_traits>
10
11#include <ddc/ddc.hpp>
12
13#include <Kokkos_Core.hpp>
14
15#if __has_include(<mkl_lapacke.h>)
16# include <mkl_lapacke.h>
17#else
18# include <lapacke.h>
19#endif
20
21// The two following headers are necessary to workaround missing includes in KokkosBatched_Getrs.hpp.
22// They must be placed before `#include <KokkosBatched_Getrs.hpp>`.
23// This is fixed in Kokkos Kernels >=5.1.
24// clang-format off
25#include <KokkosBatched_Laswp.hpp>
26#include <KokkosBatched_Trsm_Decl.hpp>
27// clang-format on
28
29#include <KokkosBatched_Getrs.hpp>
30#include <KokkosBatched_Util.hpp>
31
34
35namespace ddc::detail {
36
37template <class ExecSpace>
38SplinesLinearProblemDense<ExecSpace>::SplinesLinearProblemDense(std::size_t const mat_size)
39 : SplinesLinearProblem<ExecSpace>(mat_size)
40 , m_a("a", mat_size, mat_size)
41 , m_ipiv("ipiv", mat_size)
42{
43 Kokkos::deep_copy(m_a.view_host(), 0.);
44}
45
46template <class ExecSpace>
47SplinesLinearProblemDense<ExecSpace>::~SplinesLinearProblemDense() = default;
48
49template <class ExecSpace>
50Real SplinesLinearProblemDense<ExecSpace>::get_element(std::size_t const i, std::size_t const j)
51 const
52{
53 assert(i < size());
54 assert(j < size());
55 return m_a.view_host()(i, j);
56}
57
58template <class ExecSpace>
59void SplinesLinearProblemDense<ExecSpace>::set_element(
60 std::size_t const i,
61 std::size_t const j,
62 Real const aij)
63{
64 assert(i < size());
65 assert(j < size());
66 m_a.view_host()(i, j) = aij;
67}
68
69template <class ExecSpace>
70void SplinesLinearProblemDense<ExecSpace>::setup_solver()
71{
72 int info;
73 if constexpr (std::is_same_v<Real, float>) {
74 info = LAPACKE_sgetrf(
75 LAPACK_ROW_MAJOR,
76 size(),
77 size(),
78 m_a.view_host().data(),
79 size(),
80 m_ipiv.view_host().data());
81 } else {
82 info = LAPACKE_dgetrf(
83 LAPACK_ROW_MAJOR,
84 size(),
85 size(),
86 m_a.view_host().data(),
87 size(),
88 m_ipiv.view_host().data());
89 }
90 if (info != 0) {
91 throw std::runtime_error("LAPACKE_getrf failed with error code " + std::to_string(info));
92 }
93
94 // Convert 1-based index to 0-based index
95 for (std::size_t i = 0; i < size(); ++i) {
96 m_ipiv.view_host()(i) -= 1;
97 }
98
99 // Push on device
100 m_a.modify_host();
101 m_a.sync_device();
102 m_ipiv.modify_host();
103 m_ipiv.sync_device();
104}
105
106template <class ExecSpace>
107void SplinesLinearProblemDense<ExecSpace>::solve(MultiRHS const b, bool const transpose) const
108{
109 assert(b.extent(0) == size());
110
111 // For order 1 splines, size() can be 0 then we bypass the solver call.
112 if (size() == 0) {
113 return;
114 }
115
116 auto a_device = m_a.view_device();
117 auto ipiv_device = m_ipiv.view_device();
118
119 Kokkos::RangePolicy<ExecSpace> const policy(0, b.extent(1));
120
121 if (transpose) {
122 Kokkos::parallel_for(
123 "gerts",
124 policy,
125 KOKKOS_LAMBDA(int const i) {
126 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
127 KokkosBatched::SerialGetrs<
128 KokkosBatched::Trans::Transpose,
129 KokkosBatched::Algo::Getrs::Unblocked>::
130 invoke(a_device, ipiv_device, sub_b);
131 });
132 } else {
133 Kokkos::parallel_for(
134 "gerts",
135 policy,
136 KOKKOS_LAMBDA(int const i) {
137 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
138 KokkosBatched::SerialGetrs<
139 KokkosBatched::Trans::NoTranspose,
140 KokkosBatched::Algo::Getrs::Unblocked>::
141 invoke(a_device, ipiv_device, sub_b);
142 });
143 }
144}
145
146#if defined(KOKKOS_ENABLE_SERIAL)
147template class SplinesLinearProblemDense<Kokkos::Serial>;
148#endif
149#if defined(KOKKOS_ENABLE_OPENMP)
150template class SplinesLinearProblemDense<Kokkos::OpenMP>;
151#endif
152#if defined(KOKKOS_ENABLE_CUDA)
153template class SplinesLinearProblemDense<Kokkos::Cuda>;
154#endif
155#if defined(KOKKOS_ENABLE_HIP)
156template class SplinesLinearProblemDense<Kokkos::HIP>;
157#endif
158#if defined(KOKKOS_ENABLE_SYCL)
159template class SplinesLinearProblemDense<Kokkos::SYCL>;
160#endif
161
162} // namespace ddc::detail
The top-level namespace of DDC.