@@ -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