Tpetra parallel linear algebra Version of the Day
Loading...
Searching...
No Matches
TpetraExt_MatrixMatrix_HIP.hpp
1// @HEADER
2// *****************************************************************************
3// Tpetra: Templated Linear Algebra Services Package
4//
5// Copyright 2008 NTESS and the Tpetra contributors.
6// SPDX-License-Identifier: BSD-3-Clause
7// *****************************************************************************
8// @HEADER
9
10#ifndef TPETRA_MATRIXMATRIX_HIP_DEF_HPP
11#define TPETRA_MATRIXMATRIX_HIP_DEF_HPP
12
13#include "Tpetra_Details_IntRowPtrHelper.hpp"
14
15#ifdef HAVE_TPETRA_INST_HIP
16namespace Tpetra {
17namespace MMdetails {
18
19template <>
20struct KokkosKernelsSPGEMMBackend<Tpetra::KokkosCompat::KokkosHIPWrapperNode> {
21 static std::string parameter_prefix() { return "hip"; }
22 static std::string algorithm_label() { return "HIP"; }
23
24 template <class MatrixType>
25 static void pre_spgemm(MatrixType&) {}
26};
27
28/*********************************************************************************************************/
29// MMM KernelWrappers for Partial Specialization to HIP
30template <class Scalar,
31 class LocalOrdinal,
32 class GlobalOrdinal,
33 class LocalOrdinalViewType>
34struct KernelWrappers<Scalar, LocalOrdinal, GlobalOrdinal, Tpetra::KokkosCompat::KokkosHIPWrapperNode, LocalOrdinalViewType> {
35 using Node = Tpetra::KokkosCompat::KokkosHIPWrapperNode;
36
37 static inline void mult_A_B_newmatrix_kernel_wrapper(CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Aview,
38 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Bview,
39 const LocalOrdinalViewType& Acol2Brow,
40 const LocalOrdinalViewType& Acol2Irow,
41 const LocalOrdinalViewType& Bcol2Ccol,
42 const LocalOrdinalViewType& Icol2Ccol,
43 CrsMatrix<Scalar, LocalOrdinal, GlobalOrdinal, Node>& C,
44 Teuchos::RCP<const Import<LocalOrdinal, GlobalOrdinal, Node> > Cimport,
45 const std::string& label = std::string(),
46 const Teuchos::RCP<Teuchos::ParameterList>& params = Teuchos::null) {
47 Tpetra::MMdetails::kokkos_kernels_mult_A_B_newmatrix(
48 Aview, Bview, Acol2Brow, Acol2Irow, Bcol2Ccol, Icol2Ccol, C, Cimport, label, params);
49 }
50
51 static inline void mult_A_B_reuse_kernel_wrapper(CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Aview,
52 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Bview,
53 const LocalOrdinalViewType& Acol2Brow,
54 const LocalOrdinalViewType& Acol2Irow,
55 const LocalOrdinalViewType& Bcol2Ccol,
56 const LocalOrdinalViewType& Icol2Ccol,
57 CrsMatrix<Scalar, LocalOrdinal, GlobalOrdinal, Node>& C,
58 Teuchos::RCP<const Import<LocalOrdinal, GlobalOrdinal, Node> > Cimport,
59 const std::string& label = std::string(),
60 const Teuchos::RCP<Teuchos::ParameterList>& params = Teuchos::null) {
61 Tpetra::MMdetails::host_mult_A_B_reuse(
62 Aview, Bview, Acol2Brow, Acol2Irow,
63 Bcol2Ccol, Icol2Ccol, C, Cimport, label, params);
64 }
65};
66
67// Jacobi KernelWrappers for Partial Specialization to HIP
68template <class Scalar,
69 class LocalOrdinal,
70 class GlobalOrdinal, class LocalOrdinalViewType>
71struct KernelWrappers2<Scalar, LocalOrdinal, GlobalOrdinal, Tpetra::KokkosCompat::KokkosHIPWrapperNode, LocalOrdinalViewType> {
72 using Node = Tpetra::KokkosCompat::KokkosHIPWrapperNode;
73
74 static inline void jacobi_A_B_newmatrix_kernel_wrapper(typename Teuchos::ScalarTraits<Scalar>::magnitudeType omega,
75 const Vector<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Dinv,
76 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Aview,
77 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Bview,
78 const LocalOrdinalViewType& Acol2Brow,
79 const LocalOrdinalViewType& Acol2Irow,
80 const LocalOrdinalViewType& Bcol2Ccol,
81 const LocalOrdinalViewType& Icol2Ccol,
82 CrsMatrix<Scalar, LocalOrdinal, GlobalOrdinal, Node>& C,
83 Teuchos::RCP<const Import<LocalOrdinal, GlobalOrdinal, Node> > Cimport,
84 const std::string& label = std::string(),
85 const Teuchos::RCP<Teuchos::ParameterList>& params = Teuchos::null) {
86 // Node-specific code
87 using Teuchos::RCP;
88
89 // Options
90 // int team_work_size = 16; // Defaults to 16 as per Deveci 12/7/16 - csiefer // unreferenced
91 std::string myalg("KK");
92 if (!params.is_null()) {
93 if (params->isParameter("hip: jacobi algorithm"))
94 myalg = params->get("hip: jacobi algorithm", myalg);
95 }
96
97 if (myalg == "MSAK") {
98 ::Tpetra::MatrixMatrix::ExtraKernels::jacobi_A_B_newmatrix_MultiplyScaleAddKernel(omega, Dinv, Aview, Bview, Acol2Brow, Acol2Irow, Bcol2Ccol, Icol2Ccol, C, Cimport, label, params);
99 } else if (myalg == "KK") {
100 kokkos_kernels_jacobi_A_B_newmatrix(omega, Dinv, Aview, Bview, Acol2Brow, Acol2Irow, Bcol2Ccol, Icol2Ccol, C, Cimport, label, params);
101 } else {
102 throw std::runtime_error("Tpetra::MatrixMatrix::Jacobi newmatrix unknown kernel");
103 }
104
105 Tpetra::Details::ProfilingRegion("TpetraExt: Jacobi: Newmatrix HIPESFC");
106
107 // Final Fillcomplete
108 RCP<Teuchos::ParameterList> labelList = rcp(new Teuchos::ParameterList);
109 labelList->set("Timer Label", label);
110 if (!params.is_null()) labelList->set("compute global constants", params->get("compute global constants", true));
111
112 // NOTE: MSAK already fillCompletes, so we have to check here
113 if (!C.isFillComplete()) {
114 RCP<const Export<LocalOrdinal, GlobalOrdinal, Node> > dummyExport;
115 C.expertStaticFillComplete(Bview.origMatrix->getDomainMap(), Aview.origMatrix->getRangeMap(), Cimport, dummyExport, labelList);
116 }
117 }
118
119 static inline void jacobi_A_B_reuse_kernel_wrapper(typename Teuchos::ScalarTraits<Scalar>::magnitudeType omega,
120 const Vector<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Dinv,
121 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Aview,
122 CrsMatrixStruct<Scalar, LocalOrdinal, GlobalOrdinal, Node>& Bview,
123 const LocalOrdinalViewType& Acol2Brow,
124 const LocalOrdinalViewType& Acol2Irow,
125 const LocalOrdinalViewType& Bcol2Ccol,
126 const LocalOrdinalViewType& Icol2Ccol,
127 CrsMatrix<Scalar, LocalOrdinal, GlobalOrdinal, Node>& C,
128 Teuchos::RCP<const Import<LocalOrdinal, GlobalOrdinal, Node> > Cimport,
129 const std::string& label = std::string(),
130 const Teuchos::RCP<Teuchos::ParameterList>& params = Teuchos::null) {
131 host_jacobi_A_B_reuse(omega, Dinv, Aview, Bview, Acol2Brow, Acol2Irow,
132 Bcol2Ccol, Icol2Ccol, C, Cimport, label, params);
133 }
134};
135
136} // namespace MMdetails
137} // namespace Tpetra
138
139#endif // HIP
140
141#endif
Namespace Tpetra contains the class and methods constituting the Tpetra library.