batmat develop
Batched linear algebra routines
Loading...
Searching...
No Matches
gemm.tpp
Go to the documentation of this file.
1#pragma once
2
3#include <batmat/assume.hpp>
4// Work around GCC bug regarding declarations and explicit instantiations of extern const variables.
5#define BATMAT_LINALG_GEMM_NO_DECLARE_LUT
7#undef BATMAT_LINALG_GEMM_NO_DECLARE_LUT
9#include <batmat/loop.hpp>
10#include <batmat/ops/rotate.hpp>
11
12#define UNROLL_FOR(...) BATMAT_FULLY_UNROLLED_FOR (__VA_ARGS__)
13
15
16template <class T, class Abi, KernelConfig Conf, StorageOrder OA, StorageOrder OB, StorageOrder OC,
17 StorageOrder OD>
18BATMAT_LINALG_GEMM_EXPORT extern const constinit decltype(detail::gemm_copy_lut<T, Abi, Conf, OA,
19 OB, OC, OD>)
21
22template <MatrixStructure Struc>
23inline constexpr auto first_column =
24 [](index_t row_index) { return Struc == MatrixStructure::UpperTriangular ? row_index : 0; };
25
26template <index_t ColsReg, MatrixStructure Struc>
27inline constexpr auto last_column = [](index_t row_index) {
28 return Struc == MatrixStructure::LowerTriangular ? std::min(row_index, ColsReg - 1)
29 : ColsReg - 1;
30};
31
32/// Generalized matrix multiplication D = C ± A⁽ᵀ⁾ B⁽ᵀ⁾. Single register block.
33template <class T, class Abi, KernelConfig Conf, index_t RowsReg, index_t ColsReg, StorageOrder OA,
35[[gnu::hot, gnu::flatten]] void
37 const std::optional<uview<const T, Abi, OC>> C, const uview<T, Abi, OD> D,
38 const index_t k) noexcept {
39 static_assert(RowsReg > 0 && ColsReg > 0);
40 using enum MatrixStructure;
41 using namespace ops;
42 using simd = datapar::simd<T, Abi>;
43 // Column range for triangular matrix C (gemmt)
44 static constexpr auto min_col = first_column<Conf.struc_C>;
45 static constexpr auto max_col = last_column<ColsReg, Conf.struc_C>;
46 // The following assumption ensures that there is no unnecessary branch
47 // for k == 0 in between the loops. This is crucial for good code
48 // generation, otherwise the compiler inserts jumps and labels between
49 // the matmul kernel and the loading/storing of C, which will cause it to
50 // place C_reg on the stack, resulting in many unnecessary loads and stores.
51 BATMAT_ASSUME(k > 0);
52 // Check dimensions in the triangular case
53 if constexpr (Conf.struc_A != General)
55 if constexpr (Conf.struc_B != General)
57 // Keep C loads out of the FMA dependency chains.
58 simd C_reg[RowsReg][ColsReg]; // NOLINT(*-c-arrays)
59 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii)
60 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj)
61 C_reg[ii][jj] = simd{0};
62
63 const auto A_cached = with_cached_access<RowsReg, 0>(A);
64 const auto B_cached = with_cached_access<0, ColsReg>(B);
65
66 // Triangular matrix multiplication kernel
67 index_t l = 0;
68 if constexpr (Conf.struc_A == UpperTriangular && Conf.struc_B == LowerTriangular) {
69 l += std::max(RowsReg, ColsReg);
70 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
71 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj) {
72 UNROLL_FOR (index_t ll = std::max(ii, jj); ll < std::max(RowsReg, ColsReg); ++ll) {
73 simd &Cij = C_reg[ii][jj];
74 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, ll));
75 simd Blj = rotl<Conf.rotate_B>(B_cached.load(ll, jj));
76 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
77 }
78 }
79 }
80 } else if constexpr (Conf.struc_A == UpperTriangular) {
81 l += RowsReg;
82 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
83 UNROLL_FOR (index_t ll = ii; ll < RowsReg; ++ll) {
84 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, ll));
85 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj) {
86 simd &Cij = C_reg[ii][jj];
87 simd Blj = rotl<Conf.rotate_B>(B_cached.load(ll, jj));
88 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
89 }
90 }
91 }
92 } else if constexpr (Conf.struc_B == LowerTriangular) {
93 l += ColsReg;
94 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
95 UNROLL_FOR (index_t ll = 0; ll < ColsReg; ++ll) {
96 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, ll));
97 UNROLL_FOR (index_t jj = min_col(ii); jj <= std::min(ll, max_col(ii)); ++jj) {
98 simd &Cij = C_reg[ii][jj];
99 simd Blj = rotl<Conf.rotate_B>(B_cached.load(ll, jj));
100 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
101 }
102 }
103 }
104 }
105
106 // Rectangular matrix multiplication kernel
107 const index_t l_end_A = Conf.struc_A == LowerTriangular ? k - RowsReg : k;
108 const index_t l_end_B = Conf.struc_B == UpperTriangular ? k - ColsReg : k;
109 for (; l < std::min(l_end_A, l_end_B); ++l) {
110 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
111 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, l));
112 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj) {
113 simd &Cij = C_reg[ii][jj];
114 simd Blj = rotl<Conf.rotate_B>(B_cached.load(l, jj));
115 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
116 }
117 }
118 }
119
120 // Triangular matrix multiplication kernel
121 if constexpr (Conf.struc_A == LowerTriangular && Conf.struc_B == UpperTriangular) {
122 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
123 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj) {
124 const index_t lmax = std::min(ii, jj) + std::max(ColsReg, RowsReg) - RowsReg;
125 UNROLL_FOR (index_t ll = 0; ll <= lmax; ++ll) {
126 simd &Cij = C_reg[ii][jj];
127 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, l + ll));
128 simd Blj = rotl<Conf.rotate_B>(B_cached.load(l + ll, jj));
129 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
130 }
131 }
132 }
133 } else if constexpr (Conf.struc_A == LowerTriangular) {
134 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
135 UNROLL_FOR (index_t ll = 0; ll <= ii; ++ll) {
136 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, l + ll));
137 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj) {
138 simd &Cij = C_reg[ii][jj];
139 simd Blj = rotl<Conf.rotate_B>(B_cached.load(l + ll, jj));
140 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
141 }
142 }
143 }
144 } else if constexpr (Conf.struc_B == UpperTriangular) {
145 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii) {
146 UNROLL_FOR (index_t ll = 0; ll < ColsReg; ++ll) {
147 simd Ail = shiftl<Conf.shift_A>(A_cached.load(ii, l + ll));
148 UNROLL_FOR (index_t jj = std::max(ll, min_col(ii)); jj <= max_col(ii); ++jj) {
149 simd &Cij = C_reg[ii][jj];
150 simd Blj = rotl<Conf.rotate_B>(B_cached.load(l + ll, jj));
151 Conf.negate ? (Cij -= Ail * Blj) : (Cij += Ail * Blj);
152 }
153 }
154 }
155 }
156
157 if (C) [[likely]] {
158 const auto C_cached = with_cached_access<RowsReg, ColsReg>(*C);
159 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii)
160 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj)
161 C_reg[ii][jj] += rotl<Conf.rotate_C>(C_cached.load(ii, jj));
162 }
163
164 const auto D_cached = with_cached_access<RowsReg, ColsReg>(D);
165 // Store accumulator to memory again
166 UNROLL_FOR (index_t ii = 0; ii < RowsReg; ++ii)
167 UNROLL_FOR (index_t jj = min_col(ii); jj <= max_col(ii); ++jj)
168 D_cached.template store<Conf.mask_D>(rotr<Conf.rotate_D>(C_reg[ii][jj]), ii, jj);
169}
170
171/// Generalized matrix multiplication D = C ± A⁽ᵀ⁾ B⁽ᵀ⁾. Using register blocking.
172template <class T, class Abi, KernelConfig Conf, StorageOrder OA, StorageOrder OB, StorageOrder OC,
173 StorageOrder OD>
175 const std::optional<view<const T, Abi, OC>> C,
176 const view<T, Abi, OD> D) noexcept {
177 using enum MatrixStructure;
178 constexpr auto Rows = RowsReg<T, Abi>, Cols = ColsReg<T, Abi>;
179 // Check dimensions
180 const index_t I = D.rows(), J = D.cols(), K = A.cols();
181 BATMAT_ASSUME(A.rows() == I);
182 BATMAT_ASSUME(B.rows() == K);
183 BATMAT_ASSUME(B.cols() == J);
184 if constexpr (Conf.struc_A != General)
185 BATMAT_ASSUME(I == K);
186 if constexpr (Conf.struc_B != General)
187 BATMAT_ASSUME(K == J);
188 if constexpr (Conf.struc_C != General)
189 BATMAT_ASSUME(I == J);
190 BATMAT_ASSUME(I > 0);
191 BATMAT_ASSUME(J > 0);
192 BATMAT_ASSUME(K > 0);
193 // Configurations for the various micro-kernels
194 constexpr KernelConfig ConfGXG{.negate = Conf.negate,
195 .shift_A = Conf.shift_A,
196 .rotate_B = Conf.rotate_B,
197 .rotate_C = Conf.rotate_C,
198 .rotate_D = Conf.rotate_D,
199 .mask_D = Conf.mask_D,
200 .struc_A = General,
201 .struc_B = Conf.struc_B,
202 .struc_C = General};
203 constexpr KernelConfig ConfXGG{.negate = Conf.negate,
204 .shift_A = Conf.shift_A,
205 .rotate_B = Conf.rotate_B,
206 .rotate_C = Conf.rotate_C,
207 .rotate_D = Conf.rotate_D,
208 .mask_D = Conf.mask_D,
209 .struc_A = Conf.struc_A,
210 .struc_B = General,
211 .struc_C = General};
212 constexpr KernelConfig ConfXXG{.negate = Conf.negate,
213 .shift_A = Conf.shift_A,
214 .rotate_B = Conf.rotate_B,
215 .rotate_C = Conf.rotate_C,
216 .rotate_D = Conf.rotate_D,
217 .mask_D = Conf.mask_D,
218 .struc_A = Conf.struc_A,
219 .struc_B = Conf.struc_B,
220 .struc_C = General};
221 static const auto microkernel = gemm_copy_lut<T, Abi, Conf, OA, OB, OC, OD>;
222 static const auto microkernel_GXG = gemm_copy_lut<T, Abi, ConfGXG, OA, OB, OC, OD>;
223 static const auto microkernel_XGG = gemm_copy_lut<T, Abi, ConfXGG, OA, OB, OC, OD>;
224 static const auto microkernel_XXG = gemm_copy_lut<T, Abi, ConfXXG, OA, OB, OC, OD>;
225 // Sizeless views to partition and pass to the micro-kernels
226 const uview<const T, Abi, OA> A_ = A;
227 const uview<const T, Abi, OB> B_ = B;
228 const std::optional<uview<const T, Abi, OC>> C_ = C;
229 const uview<T, Abi, OD> D_ = D;
230
231 // Optimization for very small matrices
232 if (I <= Rows && J <= Cols)
233 return microkernel[I - 1][J - 1](A_, B_, C_, D_, K);
234
235 // Simply loop over all blocks in the given matrices.
236 auto run = [&] [[gnu::always_inline]] (index_t i, index_t ni, index_t j, index_t nj) {
237 const auto Bj = B_.middle_cols(j);
238 const auto l0A = Conf.struc_A == UpperTriangular ? i : 0;
239 const auto l1A = Conf.struc_A == LowerTriangular ? i + ni + std::max(K, I) - I : K;
240 const auto l0B = Conf.struc_B == LowerTriangular ? j : 0;
241 const auto l1B = Conf.struc_B == UpperTriangular ? j + nj + std::max(K, J) - J : K;
242 const auto l0 = std::max(l0A, l0B);
243 const auto l1 = std::min(l1A, l1B);
244 const auto Ai = A_.middle_rows(i);
245 const auto Cij = C_ ? std::make_optional(C_->block(i, j)) : std::nullopt;
246 const auto Dij = D_.block(i, j);
247 const auto Ail = Ai.middle_cols(l0);
248 const auto Blj = Bj.middle_rows(l0);
249
250 if (l1 == l0) // TODO: this is wrong.
251 return;
252 if constexpr (Conf.struc_A == LowerTriangular && Conf.struc_B == UpperTriangular) { // LU
253 if (l1A > l1B) {
254 microkernel_GXG[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
255 return;
256 } else if (l1A < l1B) {
257 microkernel_XGG[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
258 return;
259 }
260 }
261 if constexpr (Conf.struc_A == UpperTriangular && Conf.struc_B == LowerTriangular) { // UL
262 if (l0A > l0B) {
263 microkernel_XGG[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
264 return;
265 } else if (l0A < l0B) {
266 microkernel_GXG[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
267 return;
268 }
269 }
270 if constexpr (Conf.struc_C != General) { // syrk
271 if (i != j) {
272 microkernel_XXG[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
273 return;
274 }
275 }
276 microkernel[ni - 1][nj - 1](Ail, Blj, Cij, Dij, l1 - l0);
277 };
278 // Pick loop directions that allow having A=D or B=D
279 constexpr auto dir_i = Conf.struc_A == LowerTriangular ? LoopDir::Backward : LoopDir::Forward,
280 dir_j = Conf.struc_B == UpperTriangular ? LoopDir::Backward : LoopDir::Forward;
281 // Loop over block rows of A and block columns of B
282 if constexpr (OB == StorageOrder::ColMajor)
284 0, J, index_constant<Cols>(),
285 [&](index_t j, auto nj) {
286 const auto i0 = Conf.struc_C == LowerTriangular ? j : 0,
287 i1 = Conf.struc_C == UpperTriangular ? j + nj : I;
289 i0, i1, index_constant<Rows>(), [&](index_t i, auto ni) { run(i, ni, j, nj); },
290 dir_i);
291 },
292 dir_j);
293 else // swap the loops for row-major B
295 0, I, index_constant<Rows>(),
296 [&](index_t i, auto ni) {
297 const auto j0 = Conf.struc_C == UpperTriangular ? i : 0,
298 j1 = Conf.struc_C == LowerTriangular ? i + ni : J;
300 j0, j1, index_constant<Cols>(), [&](index_t j, auto nj) { run(i, ni, j, nj); },
301 dir_j);
302 },
303 dir_i);
304}
305
306} // namespace batmat::linalg::micro_kernels::gemm
#define BATMAT_ASSUME(x)
Invokes undefined behavior if the expression x does not evaluate to true.
Definition assume.hpp:17
#define UNROLL_FOR(...)
Definition gemm-diag.tpp:10
void foreach_chunked_merged(index_t i_begin, index_t i_end, auto chunk_size, auto func_chunk, LoopDir dir=LoopDir::Forward)
Iterate over the range [i_begin, i_end) in chunks of size chunk_size, calling func_chunk for each chu...
Definition loop.hpp:43
stdx::simd< Tp, Abi > simd
Definition simd.hpp:148
const constinit decltype(detail::gemm_copy_lut< T, Abi, Conf, OA, OB, OC, OD >) gemm_copy_lut
Definition gemm.tpp:20
constexpr index_t RowsReg
Register block size of the matrix-matrix multiplication micro-kernels.
Definition avx-512.hpp:13
void gemm_copy_register(view< const T, Abi, OA > A, view< const T, Abi, OB > B, std::optional< view< const T, Abi, OC > > C, view< T, Abi, OD > D) noexcept
Generalized matrix multiplication D = C ± A⁽ᵀ⁾ B⁽ᵀ⁾. Using register blocking.
Definition gemm.tpp:174
void gemm_copy_microkernel(uview< const T, Abi, OA > A, uview< const T, Abi, OB > B, std::optional< uview< const T, Abi, OC > > C, uview< T, Abi, OD > D, index_t k) noexcept
Generalized matrix multiplication D = C ± A⁽ᵀ⁾ B⁽ᵀ⁾. Single register block.
Definition gemm.tpp:36
cached_uview< Order==StorageOrder::ColMajor ? Cols :Rows, T, Abi, Order > with_cached_access(const uview< T, Abi, Order > &o) noexcept
Definition uview.hpp:228
simd_view_types< std::remove_const_t< T >, Abi >::template view< T, Order > view
Definition uview.hpp:70
std::integral_constant< index_t, I > index_constant
Definition lut.hpp:10
int index_t
Definition config.hpp:13
Self block(this const Self &self, index_t r, index_t c) noexcept
Definition uview.hpp:110
Self middle_rows(this const Self &self, index_t r) noexcept
Definition uview.hpp:114
Self middle_cols(this const Self &self, index_t c) noexcept
Definition uview.hpp:118