DDC 0.16.0
Loading...
Searching...
No Matches
splines_linear_problem_pds_band.cpp
1// Copyright (C) The DDC development team, see COPYRIGHT.md file
2//
3// SPDX-License-Identifier: MIT
4
5#include <cassert>
6#if !defined(NDEBUG)
7# include <cmath>
8# include <limits>
9#endif
10#include <cstddef>
11#include <stdexcept>
12#include <string>
13#include <type_traits>
14#include <utility>
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_Pbtrs.hpp>
27#include <KokkosBatched_Util.hpp>
28
31
32namespace ddc::detail {
33
34template <class ExecSpace>
35SplinesLinearProblemPDSBand<ExecSpace>::SplinesLinearProblemPDSBand(
36 std::size_t const mat_size,
37 std::size_t const kd)
38 : SplinesLinearProblem<ExecSpace>(mat_size)
39 , m_q("q", kd + 1, mat_size)
40{
41 assert(m_q.extent(0) <= mat_size);
42
43 Kokkos::deep_copy(m_q.view_host(), 0.);
44}
45
46template <class ExecSpace>
47SplinesLinearProblemPDSBand<ExecSpace>::~SplinesLinearProblemPDSBand() = default;
48
49template <class ExecSpace>
50Real SplinesLinearProblemPDSBand<ExecSpace>::get_element(std::size_t i, std::size_t j) const
51{
52 assert(i < size());
53 assert(j < size());
54
55 // Indices are swapped for an element on subdiagonal
56 if (i > j) {
57 std::swap(i, j);
58 }
59
60 if (j - i < m_q.extent(0)) {
61 return m_q.view_host()(j - i, i);
62 }
63
64 return 0.0;
65}
66
67template <class ExecSpace>
68void SplinesLinearProblemPDSBand<ExecSpace>::set_element(
69 std::size_t i,
70 std::size_t j,
71 Real const aij)
72{
73 assert(i < size());
74 assert(j < size());
75
76 // Indices are swapped for an element on subdiagonal
77 if (i > j) {
78 std::swap(i, j);
79 }
80 if (j - i < m_q.extent(0)) {
81 m_q.view_host()(j - i, i) = aij;
82 } else {
83 assert(std::fabs(aij) < 10 * std::numeric_limits<Real>::epsilon());
84 }
85}
86
87template <class ExecSpace>
88void SplinesLinearProblemPDSBand<ExecSpace>::setup_solver()
89{
90 int info;
91 if constexpr (std::is_same_v<Real, float>) {
92 info = LAPACKE_spbtrf(
93 LAPACK_ROW_MAJOR,
94 'L',
95 size(),
96 m_q.extent(0) - 1,
97 m_q.view_host().data(),
98 m_q.view_host().stride(
99 0) // m_q.view_host().stride(0) if LAPACK_ROW_MAJOR, m_q.view_host().stride(1) if LAPACK_COL_MAJOR
100 );
101 } else {
102 info = LAPACKE_dpbtrf(
103 LAPACK_ROW_MAJOR,
104 'L',
105 size(),
106 m_q.extent(0) - 1,
107 m_q.view_host().data(),
108 m_q.view_host().stride(
109 0) // m_q.view_host().stride(0) if LAPACK_ROW_MAJOR, m_q.view_host().stride(1) if LAPACK_COL_MAJOR
110 );
111 }
112 if (info != 0) {
113 throw std::runtime_error("LAPACKE_pbtrf failed with error code " + std::to_string(info));
114 }
115
116 // Push on device
117 m_q.modify_host();
118 m_q.sync_device();
119}
120
121template <class ExecSpace>
122void SplinesLinearProblemPDSBand<ExecSpace>::solve(MultiRHS const b, bool const) const
123{
124 assert(b.extent(0) == size());
125
126 auto q_device = m_q.view_device();
127 Kokkos::RangePolicy<ExecSpace> const policy(0, b.extent(1));
128 Kokkos::parallel_for(
129 "pbtrs",
130 policy,
131 KOKKOS_LAMBDA(int const i) {
132 auto sub_b = Kokkos::subview(b, Kokkos::ALL, i);
133 KokkosBatched::SerialPbtrs<
134 KokkosBatched::Uplo::Lower,
135 KokkosBatched::Algo::Pbtrs::Unblocked>::invoke(q_device, sub_b);
136 });
137}
138
139#if defined(KOKKOS_ENABLE_SERIAL)
140template class SplinesLinearProblemPDSBand<Kokkos::Serial>;
141#endif
142#if defined(KOKKOS_ENABLE_OPENMP)
143template class SplinesLinearProblemPDSBand<Kokkos::OpenMP>;
144#endif
145#if defined(KOKKOS_ENABLE_CUDA)
146template class SplinesLinearProblemPDSBand<Kokkos::Cuda>;
147#endif
148#if defined(KOKKOS_ENABLE_HIP)
149template class SplinesLinearProblemPDSBand<Kokkos::HIP>;
150#endif
151#if defined(KOKKOS_ENABLE_SYCL)
152template class SplinesLinearProblemPDSBand<Kokkos::SYCL>;
153#endif
154
155} // namespace ddc::detail
The top-level namespace of DDC.