Skip to content

Commit addb89e

Browse files
jan-wassenbergcopybara-github
authored andcommitted
Refactor MatMul to accept views in the kernel functions
Make arg order consistent. Move StridedView into mat.h. Add view support to RowPtrs. PiperOrigin-RevId: 804969908
1 parent f10ac41 commit addb89e

3 files changed

Lines changed: 192 additions & 146 deletions

File tree

ops/matmul-inl.h

Lines changed: 109 additions & 84 deletions
Original file line numberDiff line numberDiff line change
@@ -148,21 +148,21 @@ class MMStoreHorizontalSumsIntoC {
148148
}
149149
}
150150

151-
// Scales the dot-product terms and adds bias (if present) and stores the
152-
// four 4-wide vectors to `C` starting at `(row_c, col_c)`. If `tag` is
153-
// `MMSetC`, the vectors are written as-is (first call, or small K).
154-
// Otherwise, they are partial sums and are accumulated into C.
155-
template <class D4, class V4 = hn::Vec<D4>, class Tag, class CRows>
156-
HWY_INLINE void Store(D4 d4, V4 sum0, V4 sum1, V4 sum2, V4 sum3, Tag tag,
157-
const size_t row_c, const size_t col_c,
158-
const MMArgs& args, CRows C_rows) const {
159-
const V4 vscale = hn::Set(d4, args.scale);
151+
// Scales the dot-product terms plus `add` (if non-null) and stores the four
152+
// 4-wide vectors to `C` starting at row 0, column 0. If `tag` is `MMSetC`,
153+
// the vectors are written as-is (first call, or small K). Otherwise, they
154+
// are partial sums and are accumulated into C.
155+
template <class D4, class V4 = hn::Vec<D4>, class Tag, class CView>
156+
HWY_INLINE void Store(D4 d4, V4 sum0, V4 sum1, V4 sum2, V4 sum3,
157+
const float scale, const float* HWY_RESTRICT add,
158+
const size_t imc, Tag tag, CView C_rows) const {
159+
const V4 vscale = hn::Set(d4, scale);
160160
HWY_ALIGN static constexpr float kZero[4] = {};
161-
const V4 vadd = hn::Load(d4, args.add ? args.add + col_c : kZero);
162-
MaybeScaleAndStore<0>(d4, sum0, vscale, vadd, tag, C_rows, row_c, col_c);
163-
MaybeScaleAndStore<1>(d4, sum1, vscale, vadd, tag, C_rows, row_c, col_c);
164-
MaybeScaleAndStore<2>(d4, sum2, vscale, vadd, tag, C_rows, row_c, col_c);
165-
MaybeScaleAndStore<3>(d4, sum3, vscale, vadd, tag, C_rows, row_c, col_c);
161+
const V4 vadd = hn::Load(d4, add ? add : kZero);
162+
MaybeScaleAndStore<0>(d4, sum0, vscale, vadd, tag, imc, C_rows);
163+
MaybeScaleAndStore<1>(d4, sum1, vscale, vadd, tag, imc, C_rows);
164+
MaybeScaleAndStore<2>(d4, sum2, vscale, vadd, tag, imc, C_rows);
165+
MaybeScaleAndStore<3>(d4, sum3, vscale, vadd, tag, imc, C_rows);
166166
}
167167

168168
private:
@@ -199,13 +199,13 @@ class MMStoreHorizontalSumsIntoC {
199199
}
200200

201201
template <size_t kRow, /*deduced:*/ class DF4, class VF4 = hn::Vec<DF4>,
202-
class Tag, typename TC>
202+
class Tag, class CView>
203203
static HWY_INLINE void MaybeScaleAndStore(DF4 df4, VF4 sum, VF4 vscale,
204-
VF4 vadd, Tag, RowPtrs<TC> C_rows,
205-
const size_t row_c,
206-
const size_t col_c) {
204+
VF4 vadd, Tag, const size_t imc,
205+
CView C_view) {
207206
if constexpr (kRow < kRowsAC) {
208-
TC* HWY_RESTRICT pos = C_rows[row_c + kRow] + col_c;
207+
using TC = hwy::RemoveCvRef<decltype(C_view.Row(0)[0])>;
208+
TC* HWY_RESTRICT pos = C_view.Row(imc + kRow);
209209
const hn::Rebind<TC, DF4> dc4;
210210
if constexpr (hwy::IsSame<Tag, MMAddC>()) {
211211
vadd = F32FromTC(dc4, hn::Load(dc4, pos)); // load prior value
@@ -234,7 +234,7 @@ class MMDecompress {
234234

235235
// Neither A nor B require padding because `LoopKC` handles remainders.
236236
if constexpr (hwy::IsSame<TB, BF16>()) {
237-
return View(B, row_b, range_kc.begin(), range_kc.Num());
237+
return StridedViewBF(B, row_b, range_kc.begin(), range_kc.Num());
238238
}
239239

240240
const PackedSpan<const TB> B_span = B.PaddedSpan();
@@ -264,7 +264,7 @@ class MMDecompress {
264264
if constexpr (IsBF16<TA>()) {
265265
// We can use a view, regardless of columns/padding, because
266266
// `MMKernel::LoopKC` supports non-vector multiples.
267-
return View(A, 0, 0, A.Cols());
267+
return StridedViewBF(A, 0, 0, A.Cols());
268268
} else {
269269
// Always decompress. To reduce code size/compile time, we no longer
270270
// support a separate F32 kernel; most A are already BF16. We also only
@@ -277,15 +277,6 @@ class MMDecompress {
277277
}
278278

279279
private:
280-
// Returns 2D subrange whose top-left is `r, c` and width is `cols`.
281-
template <typename T>
282-
static StridedView<T> View(const MatPtrT<T>& AB, size_t r, size_t c,
283-
size_t cols) {
284-
HWY_DASSERT(c < AB.Cols());
285-
HWY_DASSERT(cols <= AB.Cols() - c);
286-
return StridedView<T>(const_cast<T*>(AB.Row(r)) + c, cols, AB.Stride());
287-
}
288-
289280
// Decompresses all `M x K` from `A` into padded BF16 `A_view`.
290281
static HWY_NOINLINE void DecompressA(const MatPtrT<float>& A,
291282
const StridedViewBF A_view,
@@ -402,26 +393,26 @@ class MMKernel {
402393
kMaxKC + 2 * CacheInfo::MaxLineBytes() / sizeof(BF16);
403394

404395
public:
405-
// Calls `LoopKC` for each of `mc` rows of A in steps of `mr`. `A_view`
406-
// is `mc x kc` and `B_view` is `(kNR x kc)`. Both start at row/col 0.
407396
// A2C0 in MOMMS terminology updates a `mc x kNR` slice of the output.
397+
// Calls `LoopKC` for each of `mc` rows of A in steps of `mr`. `A_view` is
398+
// `mc x kc` and `B_view` is `(kNR x kc)`. All views, including `add`, start
399+
// at row/col 0. `CView` is either `RowPtrs<TC>` or `StridedView<TC>`.
408400
// Called by B3A2C0 and by callers that hoist `A_view`.
409-
template <class Tag, class CRows>
401+
template <class Tag, class CView>
410402
static HWY_INLINE void A2C0(const StridedViewBF A_view,
411403
const StridedViewBF B_view, size_t mr,
412-
const IndexRange& range_mc, const size_t row_b,
413-
size_t kc, Tag tag, const MMArgs& args,
414-
CRows C_rows) {
404+
const IndexRange& range_mc, size_t kc,
405+
const float scale, const float* HWY_RESTRICT add,
406+
Tag tag, CView C_view) {
415407
HWY_DASSERT(1 <= mr && mr <= kMaxMR);
416-
const size_t row0 = range_mc.begin();
408+
417409
const size_t mc = range_mc.Num();
418410
size_t imc = 0;
419411

420412
// M == 1, or x86 with 8 SIMD registers:
421413
if (HWY_UNLIKELY(mr == 1)) {
422414
for (; imc < mc; ++imc) {
423-
LoopKC<1>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args,
424-
C_rows);
415+
LoopKC<1>(A_view, B_view, imc, kc, scale, add, tag, C_view);
425416
}
426417
return;
427418
}
@@ -430,32 +421,29 @@ class MMKernel {
430421
if (HWY_UNLIKELY(mr == 2)) {
431422
if (HWY_LIKELY(mc >= 2)) {
432423
for (; imc <= mc - 2; imc += 2) {
433-
LoopKC<2>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args,
434-
C_rows);
424+
LoopKC<2>(A_view, B_view, imc, kc, scale, add, tag, C_view);
435425
}
436426
}
437427
if (HWY_UNLIKELY(imc != mc)) {
438-
LoopKC<1>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args,
439-
C_rows);
428+
LoopKC<1>(A_view, B_view, imc, kc, scale, add, tag, C_view);
440429
}
441430
return;
442431
}
443432

444433
HWY_DASSERT(mr == 4);
445434
if (HWY_LIKELY(mc >= 4)) {
446435
for (; imc <= mc - 4; imc += 4) {
447-
LoopKC<4>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args,
448-
C_rows);
436+
LoopKC<4>(A_view, B_view, imc, kc, scale, add, tag, C_view);
449437
}
450438
}
451439
const size_t remainder_mc = mc - imc;
452440
HWY_DASSERT(remainder_mc < 4);
453441
if (HWY_UNLIKELY(remainder_mc & 2)) {
454-
LoopKC<2>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args, C_rows);
442+
LoopKC<2>(A_view, B_view, imc, kc, scale, add, tag, C_view);
455443
imc += 2;
456444
}
457445
if (HWY_UNLIKELY(remainder_mc & 1)) {
458-
LoopKC<1>(A_view, B_view, row0 + imc, imc, row_b, kc, tag, args, C_rows);
446+
LoopKC<1>(A_view, B_view, imc, kc, scale, add, tag, C_view);
459447
imc += 1;
460448
}
461449
HWY_DASSERT(imc == mc);
@@ -466,11 +454,11 @@ class MMKernel {
466454
// Loop over NC/MC/KC, called from the outer loops. The MOMMS B3A2C0 reads
467455
// `mc x kc` of A, `nc x kc` of B, and updates `mc x nc` of C. Called by
468456
// `ForeachKC` and when there is only a single KC task.
469-
template <typename TB, typename Tag, class CRows>
457+
template <typename TB, typename TC, typename Tag>
470458
static void B3A2C0(const StridedViewBF A, const MatPtrT<TB>& B,
471-
const MMArgs& args, const IndexRange& range_mc,
472-
const IndexRange& range_kc, const IndexRange& range_nc,
473-
size_t mr, Tag out_tag, CRows C_rows) {
459+
const IndexRange& range_mc, const IndexRange& range_kc,
460+
const IndexRange& range_nc, const MMArgs& args,
461+
Tag out_tag, RowPtrs<TC> C) {
474462
HWY_ALIGN BF16 B_storage[B_storage_max];
475463

476464
const size_t kc = range_kc.Num();
@@ -482,24 +470,28 @@ class MMKernel {
482470

483471
for (size_t row_b = range_nc.begin(); row_b < range_nc.end();
484472
row_b += kNR) {
485-
StridedViewBF B_view =
473+
const StridedViewBF B_view =
486474
MMDecompress::DecompressB(B, row_b, range_kc, B_storage_view);
487-
A2C0(A_view, B_view, mr, range_mc, row_b, kc, out_tag, args, C_rows);
475+
const RowPtrs<TC> C_view = C.View(range_mc.begin(), row_b);
476+
const float* HWY_RESTRICT add = args.add ? args.add + row_b : nullptr;
477+
A2C0(A_view, B_view, args.mr, range_mc, kc, args.scale, add, out_tag,
478+
C_view);
488479
}
489480
}
490481

491-
template <typename TB, class CRows>
482+
template <typename TB, typename TC>
492483
static void ForeachKC(const StridedViewBF A, const MatPtrT<TB>& B,
493-
const MMArgs& args, const IndexRange& range_mc,
484+
const IndexRange& range_mc,
494485
const IndexRangePartition& ranges_kc,
495-
const IndexRange& range_nc, size_t mr, CRows C_rows) {
486+
const IndexRange& range_nc, const MMArgs& args,
487+
RowPtrs<TC> C) {
496488
// Peel off the first iteration of the kc loop: avoid zero-initializing `C`
497489
// by writing directly into it, and later accumulating into it.
498490
ranges_kc.VisitFirst([&](const IndexRange& range_kc) {
499-
B3A2C0(A, B, args, range_mc, range_kc, range_nc, mr, MMSetC(), C_rows);
491+
B3A2C0(A, B, range_mc, range_kc, range_nc, args, MMSetC(), C);
500492
});
501493
ranges_kc.VisitRemaining([&](const IndexRange& range_kc) {
502-
B3A2C0(A, B, args, range_mc, range_kc, range_nc, mr, MMAddC(), C_rows);
494+
B3A2C0(A, B, range_mc, range_kc, range_nc, args, MMAddC(), C);
503495
});
504496
}
505497

@@ -593,19 +585,20 @@ class MMKernel {
593585
// Innermost loop over `kc` columns (typically 1024-4096, not necessarily a
594586
// multiple of `NBF`) in steps of one vector, for `kRowsAC` rows of `A_view`
595587
// from range_mc-relative `imc` and `B_view` from row 0 (both at column 0).
596-
// Updates a `kRowsAC x kNR` tile with top-left `C.Row(row_ac) + col_c`.
597-
// `A` and `B` are always BF16, `C` can be F32 or BF16.
598-
template <size_t kRowsAC, /*deduced:*/ class Tag, class CRows>
588+
// Updates a `kRowsAC x kNR` tile in `C_view` starting at row `imc`, column 0.
589+
// `A` and `B` are always BF16, `C` can be F32 or BF16. `add` is also
590+
// relative to the C column.
591+
template <size_t kRowsAC, /*deduced:*/ class Tag, class CView>
599592
static HWY_INLINE void LoopKC(const StridedViewBF A_view,
600-
const StridedViewBF B_view, size_t row_ac,
601-
size_t imc, size_t col_c, size_t kc, Tag tag,
602-
const MMArgs& args, CRows C_rows) {
593+
const StridedViewBF B_view, size_t imc,
594+
size_t kc, const float scale,
595+
const float* HWY_RESTRICT add, Tag tag,
596+
CView C_view) {
603597
const hn::ScalableTag<BF16> dbf;
604598
using VBF = hn::Vec<decltype(dbf)>;
605599
HWY_LANES_CONSTEXPR const size_t NBF = hn::Lanes(dbf);
606600

607601
HWY_DASSERT(kRowsAC <= kMaxMR);
608-
HWY_DASSERT(col_c % kNR == 0);
609602
// Rows are aligned to `kMaxMR`, except for the last tile of A.
610603

611604
// `kRowsAC` rows of A (null for the rest) and `kNR` rows of B.
@@ -784,7 +777,7 @@ class MMKernel {
784777
hn::Vec<decltype(d4)> sum0, sum1, sum2, sum3;
785778
horz.Reduce4x4(df, C00, C01, C02, C03, C10, C11, C12, C13, C20, C21, C22,
786779
C23, C30, C31, C32, C33, sum0, sum1, sum2, sum3);
787-
horz.Store(d4, sum0, sum1, sum2, sum3, tag, row_ac, col_c, args, C_rows);
780+
horz.Store(d4, sum0, sum1, sum2, sum3, scale, add, imc, tag, C_view);
788781
}
789782
};
790783

@@ -884,15 +877,15 @@ class MMLoops {
884877
// or with the best config.
885878
template <typename TB, typename TC>
886879
static HWY_NOINLINE void Dispatch(const StridedViewBF A, const MatPtrT<TB>& B,
887-
RowPtrs<TC> C_rows, const MMArgs& args) {
880+
RowPtrs<TC> C, const MMArgs& args) {
888881
static const auto zone = args.env.ctx.profiler.AddZone("MM.Dispatch");
889882
PROFILER_ZONE3(args.env.ctx.profiler,
890883
args.env.ctx.Worker(args.options.cluster_idx), zone);
891884

892885
DispatchParallelism(
893886
args.options.parallelism, [&](const auto& parallel) HWY_ATTR {
894887
DispatchOrder(args.order, [&](const auto& order) HWY_ATTR {
895-
Loop(order, parallel, A, B, C_rows, args);
888+
Loop(order, parallel, A, B, C, args);
896889
});
897890
});
898891
}
@@ -904,11 +897,11 @@ class MMLoops {
904897
return HWY_MAX(kNR, line_bytes / sizeof_TC);
905898
}
906899

907-
// Single M and K ranges, parallel N. Fills all of C directly.
900+
// Single M and K ranges, parallel N.
908901
template <typename TB, typename TC, class Parallel>
909902
static HWY_INLINE void Loop(MMOrderNT, Parallel parallel,
910903
const StridedViewBF A, const MatPtrT<TB>& B,
911-
RowPtrs<TC> C_rows, const MMArgs& args) {
904+
RowPtrs<TC> C, const MMArgs& args) {
912905
static const auto zone = args.env.ctx.profiler.AddZone("MM.NT");
913906
HWY_DASSERT(args.ranges_mc.NumTasks() == 1);
914907
HWY_DASSERT(args.ranges_kc.NumTasks() == 1);
@@ -932,10 +925,21 @@ class MMLoops {
932925

933926
for (size_t row_b = range_nc.begin(); row_b < range_nc.end();
934927
row_b += kNR) {
935-
StridedViewBF B_view =
928+
const StridedViewBF B_view =
936929
MMDecompress::DecompressB(B, row_b, range_K, B_storage_view);
937-
MMKernel::A2C0(A_view, B_view, args.mr, range_M, row_b, K, MMSetC(),
938-
args, C_rows);
930+
const RowPtrs<TC> C_view = C.View(range_M.begin(), row_b);
931+
const float* HWY_RESTRICT add =
932+
args.add ? args.add + row_b : nullptr;
933+
934+
MMKernel::A2C0(A_view, B_view, args.mr, range_M, K, args.scale, add,
935+
MMSetC(), C_view);
936+
}
937+
938+
if constexpr (IsBF16<TC>()) {
939+
if (args.options.fused) {
940+
StridedViewBF C2(nullptr, 0, 0);
941+
args.options.fused(C, range_M, range_nc, C2, worker);
942+
}
939943
}
940944
});
941945
}
@@ -944,7 +948,7 @@ class MMLoops {
944948
template <typename TB, typename TC, class Parallel>
945949
static HWY_INLINE void Loop(MMOrderNT_K, Parallel parallel,
946950
const StridedViewBF A, const MatPtrT<TB>& B,
947-
RowPtrs<TC> C_rows, const MMArgs& args) {
951+
RowPtrs<TC> C, const MMArgs& args) {
948952
static const auto zone = args.env.ctx.profiler.AddZone("MM.NT_K");
949953
HWY_DASSERT(args.ranges_mc.NumTasks() == 1);
950954
const IndexRange& range_mc = args.ranges_mc.Range(0);
@@ -955,17 +959,24 @@ class MMLoops {
955959
[&](const IndexRange& range_nc, size_t worker) HWY_ATTR {
956960
MMZone mm_zone;
957961
mm_zone.MaybeEnter(worker, zone, args.env, &args.autotune);
958-
MMKernel::ForeachKC(A, B, args, range_mc, args.ranges_kc,
959-
range_nc, args.mr, C_rows);
962+
MMKernel::ForeachKC(A, B, range_mc, args.ranges_kc,
963+
range_nc, args, C);
964+
965+
if constexpr (IsBF16<TC>()) {
966+
if (args.options.fused) {
967+
StridedViewBF C2(nullptr, 0, 0);
968+
args.options.fused(C, range_mc, range_nc, C2, worker);
969+
}
970+
}
960971
});
961972
}
962973

963974
// Parallel loops over mc/nc blocks of M/range_n, single K.
964-
// Fills `mc x nc` sections of C directly, in parallel.
975+
// Fills `mc x nc` sections of C.
965976
template <typename TB, typename TC, class Parallel>
966977
static HWY_INLINE void Loop(MMOrderNT_MT, Parallel parallel,
967978
const StridedViewBF A, const MatPtrT<TB>& B,
968-
RowPtrs<TC> C_rows, const MMArgs& args) {
979+
RowPtrs<TC> C, const MMArgs& args) {
969980
static const auto zone = args.env.ctx.profiler.AddZone("MM.NT_MT");
970981
HWY_DASSERT(args.ranges_kc.NumTasks() == 1);
971982
const IndexRange& range_K = args.ranges_kc.Range(0);
@@ -976,17 +987,24 @@ class MMLoops {
976987
size_t worker) HWY_ATTR {
977988
MMZone mm_zone;
978989
mm_zone.MaybeEnter(worker, zone, args.env, &args.autotune);
979-
MMKernel::B3A2C0(A, B, args, range_mc, range_K, range_nc, args.mr,
980-
MMSetC(), C_rows);
990+
MMKernel::B3A2C0(A, B, range_mc, range_K, range_nc, args, MMSetC(),
991+
C);
992+
993+
if constexpr (IsBF16<TC>()) {
994+
if (args.options.fused) {
995+
StridedViewBF C2(nullptr, 0, 0);
996+
args.options.fused(C, range_mc, range_nc, C2, worker);
997+
}
998+
}
981999
});
9821000
}
9831001

984-
// Parallel loops over mc/nc blocks of M/range_np, sequential K.
1002+
// Parallel loops over mc/nc blocks of M/range_n, sequential K.
9851003
// Accumulates into `mc x nc` sections of `C`.
9861004
template <typename TB, typename TC, class Parallel>
9871005
static HWY_INLINE void Loop(MMOrderNT_MT_K, Parallel parallel,
9881006
const StridedViewBF A, const MatPtrT<TB>& B,
989-
RowPtrs<TC> C_rows, const MMArgs& args) {
1007+
RowPtrs<TC> C, const MMArgs& args) {
9901008
static const auto zone = args.env.ctx.profiler.AddZone("MM.NT_MT_K");
9911009

9921010
parallel.ForRangesMC_NC(
@@ -995,8 +1013,15 @@ class MMLoops {
9951013
size_t worker) HWY_ATTR {
9961014
MMZone mm_zone;
9971015
mm_zone.MaybeEnter(worker, zone, args.env, &args.autotune);
998-
MMKernel::ForeachKC(A, B, args, range_mc, args.ranges_kc, range_nc,
999-
args.mr, C_rows);
1016+
MMKernel::ForeachKC(A, B, range_mc, args.ranges_kc, range_nc, args,
1017+
C);
1018+
1019+
if constexpr (IsBF16<TC>()) {
1020+
if (args.options.fused) {
1021+
StridedViewBF C2(nullptr, 0, 0);
1022+
args.options.fused(C, range_mc, range_nc, C2, worker);
1023+
}
1024+
}
10001025
});
10011026
}
10021027
}; // MMLoops

0 commit comments

Comments
 (0)