mirror of
https://github.com/opencv/opencv.git
synced 2026-07-21 19:33:03 +04:00
42f41f749a
3rdparty(mlas): add missing power (ppc64le) kernel headers #29516 Closes: https://github.com/opencv/opencv/issues/29465 ### Pull Request Readiness Checklist See details at https://github.com/opencv/opencv/wiki/How_to_contribute#making-a-good-pull-request - [x] I agree to contribute to the project under Apache 2 License. - [x] To the best of my knowledge, the proposed patch is not based on a code under GPL or another license that is incompatible with OpenCV - [x] The PR is proposed to the proper branch - [x] There is a reference to the original bug report and related work - [x] There is accuracy test, performance test and test data in opencv_extra repository, if applicable Patch to opencv_extra has the same branch name. - [x] The feature is well documented and sample code can be built with the project CMake
699 lines
28 KiB
C++
699 lines
28 KiB
C++
/*++
|
|
|
|
Copyright (c) Microsoft Corporation. All rights reserved.
|
|
|
|
Licensed under the MIT License.
|
|
|
|
Module Name:
|
|
|
|
SgemmKernelPower.cpp
|
|
|
|
Abstract:
|
|
|
|
This module implements the kernels for the single precision matrix/matrix
|
|
multiply operation (SGEMM).
|
|
|
|
--*/
|
|
|
|
#define PREFETCH_ADDR(addr) \
|
|
asm volatile("dcbt 0, %0" ::"r"(addr) : "memory");
|
|
|
|
#include "SgemmKernelpower.h"
|
|
#include <cstring> // memset; OpenCV's vendored mlasi.h subset does not pull in <cstring>
|
|
extern "C" void
|
|
PackAKernelPOWER10(__vector float* D, const float* A, size_t lda, size_t k, size_t RowCount);
|
|
struct MlasSgemmBroadcastAElementsMMA
|
|
{
|
|
template<size_t RowCount, size_t Row>
|
|
MLAS_FORCEINLINE
|
|
static
|
|
void
|
|
Iteration(
|
|
MLAS_FLOAT32X4 ABroadcast[RowCount],
|
|
const float* A,
|
|
size_t lda
|
|
)
|
|
{
|
|
ABroadcast[0] = vec_insert(A[Row * lda], ABroadcast[0], Row);
|
|
}
|
|
};
|
|
|
|
template<size_t RowCount>
|
|
MLAS_FORCEINLINE
|
|
void
|
|
MlasSgemmComputeAElements(
|
|
MLAS_FLOAT32X4 AElements[RowCount],
|
|
MLAS_FLOAT32X4 ABroadcast[RowCount]
|
|
)
|
|
{
|
|
__vector float a1,a2;
|
|
a1 = vec_mergee (AElements[0], AElements[1]);
|
|
a2 = vec_mergee (AElements[2], AElements[3]);
|
|
ABroadcast[0] =vec_xxpermdi(a1,a2,0);
|
|
ABroadcast[2] =vec_xxpermdi(a1,a2,3);
|
|
a1 = vec_mergeo (AElements[0], AElements[1]);
|
|
a2 = vec_mergeo (AElements[2], AElements[3]);
|
|
ABroadcast[1] =vec_xxpermdi(a1,a2,0);
|
|
ABroadcast[3] =vec_xxpermdi(a1,a2,3);
|
|
}
|
|
template<size_t RowCount>
|
|
MLAS_FORCEINLINE
|
|
void
|
|
MlasSgemmComputeBlockMMA(
|
|
__vector_quad acc[8],
|
|
MLAS_FLOAT32X4 ABroadcast,
|
|
MLAS_FLOAT32X4 A2Broadcast,
|
|
const float* B,
|
|
size_t CountM
|
|
)
|
|
{
|
|
MLAS_FLOAT32X4 BElements[4];
|
|
typedef __vector unsigned char vec_t;
|
|
|
|
BElements[0] = MlasLoadFloat32x4(B);
|
|
BElements[1] = MlasLoadFloat32x4(B + 4);
|
|
BElements[2] = MlasLoadFloat32x4(B + 8);
|
|
BElements[3] = MlasLoadFloat32x4(B + 12);
|
|
__builtin_mma_xvf32gerpp (&acc[0], reinterpret_cast<vec_t>(ABroadcast), reinterpret_cast<vec_t>(BElements[0]));
|
|
__builtin_mma_xvf32gerpp (&acc[1], reinterpret_cast<vec_t>(ABroadcast), reinterpret_cast<vec_t>(BElements[1]));
|
|
__builtin_mma_xvf32gerpp (&acc[2], reinterpret_cast<vec_t>(ABroadcast), reinterpret_cast<vec_t>(BElements[2]));
|
|
__builtin_mma_xvf32gerpp (&acc[3], reinterpret_cast<vec_t>(ABroadcast), reinterpret_cast<vec_t>(BElements[3]));
|
|
if (CountM == 8) {
|
|
__builtin_mma_xvf32gerpp (&acc[4], reinterpret_cast<vec_t>(A2Broadcast), reinterpret_cast<vec_t>(BElements[0]));
|
|
__builtin_mma_xvf32gerpp (&acc[5], reinterpret_cast<vec_t>(A2Broadcast), reinterpret_cast<vec_t>(BElements[1]));
|
|
__builtin_mma_xvf32gerpp (&acc[6], reinterpret_cast<vec_t>(A2Broadcast), reinterpret_cast<vec_t>(BElements[2]));
|
|
__builtin_mma_xvf32gerpp (&acc[7], reinterpret_cast<vec_t>(A2Broadcast), reinterpret_cast<vec_t>(BElements[3]));
|
|
}
|
|
}
|
|
template<size_t VectorCount>
|
|
struct MlasSgemmStoreVectorMMA
|
|
{
|
|
template<size_t RowCount, size_t Row>
|
|
MLAS_FORCEINLINE
|
|
static
|
|
void
|
|
Iteration(
|
|
MLAS_FLOAT32X4 Result[4],
|
|
float* C,
|
|
size_t ldc,
|
|
MLAS_FLOAT32X4 AlphaBroadcast,
|
|
bool ZeroMode
|
|
)
|
|
{
|
|
MLAS_FLOAT32X4 *rowC;
|
|
if (ZeroMode) {
|
|
rowC = reinterpret_cast<MLAS_FLOAT32X4 *>(&C[Row * ldc + VectorCount]);
|
|
rowC[0] = Result[Row] * AlphaBroadcast;
|
|
} else {
|
|
rowC = reinterpret_cast<MLAS_FLOAT32X4 *>(&C[Row * ldc + VectorCount]);
|
|
rowC[0] += Result[Row] * AlphaBroadcast;
|
|
}
|
|
}
|
|
};
|
|
|
|
struct MlasSgemmMultiplyAlphaTrailingMMA
|
|
{
|
|
template<size_t RowCount, size_t Row>
|
|
MLAS_FORCEINLINE
|
|
static
|
|
void
|
|
Iteration(
|
|
MLAS_FLOAT32X4 Accumulators[RowCount],
|
|
MLAS_FLOAT32X4 AlphaBroadcast
|
|
)
|
|
{
|
|
Accumulators[Row] = MlasMultiplyFloat32x4(Accumulators[Row], AlphaBroadcast);
|
|
}
|
|
};
|
|
template<unsigned Lane>
|
|
struct MlasSgemmStoreScalarMMA
|
|
{
|
|
template<size_t RowCount, size_t Row>
|
|
MLAS_FORCEINLINE
|
|
static
|
|
void
|
|
Iteration(
|
|
MLAS_FLOAT32X4 Accumulators[RowCount],
|
|
float* C,
|
|
size_t ldc,
|
|
bool ZeroMode
|
|
)
|
|
{
|
|
float* c = C + Row * ldc + Lane;
|
|
float Value = Accumulators[Row][Lane];
|
|
if (!ZeroMode) {
|
|
Value += *c;
|
|
}
|
|
|
|
*c = Value;
|
|
}
|
|
};
|
|
|
|
template <size_t RowCount>
|
|
MLAS_FORCEINLINE
|
|
size_t
|
|
MlasSgemmMMAProcessCount(
|
|
__vector float* Pa,
|
|
const float* B,
|
|
float* C,
|
|
size_t CountM,
|
|
size_t CountK,
|
|
size_t CountN,
|
|
size_t ldc,
|
|
MLAS_FLOAT32X4 AlphaBroadcast,
|
|
bool ZeroMode
|
|
)
|
|
{
|
|
do {
|
|
__vector float* pa1 = Pa;
|
|
size_t k = CountK;
|
|
|
|
MLAS_FLOAT32X4 Accumulators[2][RowCount] = {{0}};
|
|
MLAS_FLOAT32X4 Result[RowCount];
|
|
MLAS_FLOAT32X4 ABroadcast[RowCount] = {0};
|
|
__vector_quad acc[8];
|
|
|
|
//
|
|
// Clear the block accumulators.
|
|
//
|
|
__builtin_mma_xxsetaccz(&acc[0]);
|
|
__builtin_mma_xxsetaccz(&acc[1]);
|
|
__builtin_mma_xxsetaccz(&acc[2]);
|
|
__builtin_mma_xxsetaccz(&acc[3]);
|
|
__builtin_mma_xxsetaccz(&acc[4]);
|
|
__builtin_mma_xxsetaccz(&acc[5]);
|
|
__builtin_mma_xxsetaccz(&acc[6]);
|
|
__builtin_mma_xxsetaccz(&acc[7]);
|
|
|
|
//
|
|
// Compute the output block.
|
|
//
|
|
while (k >= 8) {
|
|
if (CountM == 8) {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], pa1[4], B, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[1], pa1[5], B + 16, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[2], pa1[6], B + 32, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[3], pa1[7], B + 48, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[8], pa1[12], B + 64, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[9], pa1[13], B + 80, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[10], pa1[14], B + 96, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[11], pa1[15], B + 112, CountM);
|
|
B += 128;
|
|
pa1 += 16;
|
|
k -= 8;
|
|
} else {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], ABroadcast[0], B, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[1], ABroadcast[1], B + 16, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[2], ABroadcast[2], B + 32, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[3], ABroadcast[3], B + 48, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[4], ABroadcast[0], B + 64, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[5], ABroadcast[1], B + 80, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[6], ABroadcast[2], B + 96, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[7], ABroadcast[3], B + 112, CountM);
|
|
B += 128;
|
|
pa1 += 8;
|
|
k -= 8;
|
|
}
|
|
}
|
|
|
|
while (k >= 4) {
|
|
if (CountM == 8) {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], pa1[4], B, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[1], pa1[5], B + 16, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[2], pa1[6], B + 32, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[3], pa1[7], B + 48, CountM);
|
|
B += 16 * 4;
|
|
pa1 += 8;
|
|
k -= 4;
|
|
} else {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], ABroadcast[0], B, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[1], ABroadcast[1], B + 16, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[2], ABroadcast[2], B + 32, CountM);
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[3], ABroadcast[3], B + 48, CountM);
|
|
B += 16 * 4;
|
|
pa1 += 4;
|
|
k -= 4;
|
|
}
|
|
}
|
|
while (k > 0) {
|
|
if (CountM == 8) {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], pa1[1], B, CountM);
|
|
pa1 += 2;
|
|
} else {
|
|
MlasSgemmComputeBlockMMA<RowCount>(&acc[0], pa1[0], ABroadcast[0], B, CountM);
|
|
pa1 += 1;
|
|
}
|
|
B += 16;
|
|
k -= 1;
|
|
}
|
|
if (CountN >= 16) {
|
|
|
|
//
|
|
// Store the entire output block.
|
|
//
|
|
__builtin_mma_disassemble_acc (Result, &acc[0]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[1]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[2]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<8>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[3]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<12>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountM == 8) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[4]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[5]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[6]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<8>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[7]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<12>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
}
|
|
} else {
|
|
|
|
//
|
|
// Store the partial output block.
|
|
//
|
|
|
|
if (CountN >= 12) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[0]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[1]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[2]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<8>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountM == 8) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[4]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[5]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[6]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<8>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountN - 12 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[1], &acc[7]);
|
|
}
|
|
}
|
|
if (CountN - 12 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[0], &acc[3]);
|
|
}
|
|
} else if (CountN >= 8) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[0]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[1]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountM == 8) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[4]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
__builtin_mma_disassemble_acc (Result, &acc[5]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<4>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountN - 8 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[1], &acc[6]);
|
|
}
|
|
}
|
|
if (CountN - 8 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[0], &acc[2]);
|
|
}
|
|
} else if (CountN >= 4) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[0]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C, ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountM == 8) {
|
|
__builtin_mma_disassemble_acc (Result, &acc[4]);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreVectorMMA<0>>()(Result, C + (ldc*4), ldc, AlphaBroadcast, ZeroMode);
|
|
if (CountN - 4 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[1], &acc[5]);
|
|
}
|
|
}
|
|
if (CountN - 4 > 0) {
|
|
__builtin_mma_disassemble_acc (Accumulators[0], &acc[1]);
|
|
}
|
|
} else {
|
|
__builtin_mma_disassemble_acc (Accumulators[0], &acc[0]);
|
|
if (CountM == 8) {
|
|
__builtin_mma_disassemble_acc (Accumulators[1], &acc[4]);
|
|
}
|
|
}
|
|
|
|
//
|
|
// Store the remaining unaligned columns.
|
|
//
|
|
|
|
C += (CountN & ~3);
|
|
CountN &= 3;
|
|
|
|
if (CountN > 0) {
|
|
|
|
MlasLoopUnroll<RowCount, MlasSgemmMultiplyAlphaTrailingMMA>()(Accumulators[0], AlphaBroadcast);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<0>>()(Accumulators[0], C, ldc, ZeroMode);
|
|
if (CountM == 8) {
|
|
MlasLoopUnroll<RowCount, MlasSgemmMultiplyAlphaTrailingMMA>()(Accumulators[1], AlphaBroadcast);
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<0>>()(Accumulators[1], C + (ldc*4), ldc, ZeroMode);
|
|
}
|
|
if (CountN >= 2) {
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<1>>()(Accumulators[0], C, ldc, ZeroMode);
|
|
if (CountM == 8) {
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<1>>()(Accumulators[1], C + (ldc*4), ldc, ZeroMode);
|
|
}
|
|
}
|
|
if (CountN >= 3) {
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<2>>()(Accumulators[0], C, ldc, ZeroMode);
|
|
if (CountM == 8) {
|
|
MlasLoopUnroll<RowCount, MlasSgemmStoreScalarMMA<2>>()(Accumulators[1], C + (ldc*4), ldc, ZeroMode);
|
|
}
|
|
}
|
|
}
|
|
|
|
break;
|
|
}
|
|
|
|
C += 16;
|
|
CountN -= 16;
|
|
|
|
} while (CountN > 0);
|
|
|
|
return CountM;
|
|
}
|
|
|
|
template <size_t RowCount>
|
|
MLAS_FORCEINLINE void
|
|
MlasSgemmPackA(
|
|
__vector float* D,
|
|
const float* A,
|
|
size_t lda,
|
|
size_t k
|
|
)
|
|
{
|
|
__vector float a1, a2;
|
|
const float* a = A;
|
|
MLAS_FLOAT32X4 AElements[RowCount] = {};
|
|
MLAS_FLOAT32X4 A2Elements[RowCount] = {};
|
|
while (k >= 16)
|
|
|
|
{
|
|
PREFETCH_ADDR(a);
|
|
PREFETCH_ADDR(a + lda);
|
|
PREFETCH_ADDR(a + 2 * lda);
|
|
PREFETCH_ADDR(a + 3 * lda);
|
|
PREFETCH_ADDR(a + 4 * lda);
|
|
PREFETCH_ADDR(a + 5 * lda);
|
|
PREFETCH_ADDR(a + 6 * lda);
|
|
PREFETCH_ADDR(a + 7 * lda);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[0] = vec_xxpermdi(a1, a2, 0);
|
|
D[2] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[1] = vec_xxpermdi(a1, a2, 0);
|
|
D[3] = vec_xxpermdi(a1, a2, 3);
|
|
if (RowCount == 8) {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 4, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[8] = vec_xxpermdi(a1, a2, 0);
|
|
D[10] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[9] = vec_xxpermdi(a1, a2, 0);
|
|
D[11] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 8, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[16] = vec_xxpermdi(a1, a2, 0);
|
|
D[18] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[17] = vec_xxpermdi(a1, a2, 0);
|
|
D[19] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 12, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[24] = vec_xxpermdi(a1, a2, 0);
|
|
D[26] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[25] = vec_xxpermdi(a1, a2, 0);
|
|
D[27] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, a + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[4] = vec_xxpermdi(a1, a2, 0);
|
|
D[6] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[5] = vec_xxpermdi(a1, a2, 0);
|
|
D[7] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, (a + 4) + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[12] = vec_xxpermdi(a1, a2, 0);
|
|
D[14] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[13] = vec_xxpermdi(a1, a2, 0);
|
|
D[15] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, (a + 8) + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[20] = vec_xxpermdi(a1, a2, 0);
|
|
D[22] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[21] = vec_xxpermdi(a1, a2, 0);
|
|
D[23] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, (a + 12) + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[28] = vec_xxpermdi(a1, a2, 0);
|
|
D[30] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[29] = vec_xxpermdi(a1, a2, 0);
|
|
D[31] = vec_xxpermdi(a1, a2, 3);
|
|
D += 32;
|
|
} else {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 4, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[4] = vec_xxpermdi(a1, a2, 0);
|
|
D[6] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[5] = vec_xxpermdi(a1, a2, 0);
|
|
D[7] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 8, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[8] = vec_xxpermdi(a1, a2, 0);
|
|
D[10] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[9] = vec_xxpermdi(a1, a2, 0);
|
|
D[11] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 12, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[12] = vec_xxpermdi(a1, a2, 0);
|
|
D[14] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[13] = vec_xxpermdi(a1, a2, 0);
|
|
D[15] = vec_xxpermdi(a1, a2, 3);
|
|
D += 16;
|
|
}
|
|
k -= 16;
|
|
a += 16;
|
|
}
|
|
|
|
while (k >= 8) {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[0] = vec_xxpermdi(a1, a2, 0);
|
|
D[2] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[1] = vec_xxpermdi(a1, a2, 0);
|
|
D[3] = vec_xxpermdi(a1, a2, 3);
|
|
if (RowCount == 8) {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 4, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[8] = vec_xxpermdi(a1, a2, 0);
|
|
D[10] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[9] = vec_xxpermdi(a1, a2, 0);
|
|
D[11] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, a + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[4] = vec_xxpermdi(a1, a2, 0);
|
|
D[6] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[5] = vec_xxpermdi(a1, a2, 0);
|
|
D[7] = vec_xxpermdi(a1, a2, 3);
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, (a + 4) + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[12] = vec_xxpermdi(a1, a2, 0);
|
|
D[14] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[13] = vec_xxpermdi(a1, a2, 0);
|
|
D[15] = vec_xxpermdi(a1, a2, 3);
|
|
D += 16;
|
|
} else {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a + 4, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[4] = vec_xxpermdi(a1, a2, 0);
|
|
D[6] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[5] = vec_xxpermdi(a1, a2, 0);
|
|
D[7] = vec_xxpermdi(a1, a2, 3);
|
|
D += 8;
|
|
}
|
|
a += 8;
|
|
k -= 8;
|
|
}
|
|
|
|
while (k >= 4) {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(AElements, a, lda);
|
|
a1 = vec_mergee(AElements[0], AElements[1]);
|
|
a2 = vec_mergee(AElements[2], AElements[3]);
|
|
D[0] = vec_xxpermdi(a1, a2, 0);
|
|
D[2] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(AElements[0], AElements[1]);
|
|
a2 = vec_mergeo(AElements[2], AElements[3]);
|
|
D[1] = vec_xxpermdi(a1, a2, 0);
|
|
D[3] = vec_xxpermdi(a1, a2, 3);
|
|
if (RowCount == 8) {
|
|
MlasLoopUnroll<4, MlasFgemmLoadAElements>()(A2Elements, a + (lda * 4), lda);
|
|
a1 = vec_mergee(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergee(A2Elements[2], A2Elements[3]);
|
|
D[4] = vec_xxpermdi(a1, a2, 0);
|
|
D[6] = vec_xxpermdi(a1, a2, 3);
|
|
a1 = vec_mergeo(A2Elements[0], A2Elements[1]);
|
|
a2 = vec_mergeo(A2Elements[2], A2Elements[3]);
|
|
D[5] = vec_xxpermdi(a1, a2, 0);
|
|
D[7] = vec_xxpermdi(a1, a2, 3);
|
|
D += 8;
|
|
} else
|
|
D += 4;
|
|
a += 4;
|
|
k -= 4;
|
|
}
|
|
|
|
/* When k is less than 4, copy a single element from each row. */
|
|
while (k > 0) {
|
|
MlasLoopUnroll<4, MlasSgemmBroadcastAElementsMMA>()(AElements, a, lda);
|
|
D[0] = AElements[0];
|
|
if (RowCount == 8) {
|
|
MlasLoopUnroll<4, MlasSgemmBroadcastAElementsMMA>()(A2Elements, a + (lda * 4), lda);
|
|
D[1] = A2Elements[0];
|
|
D += 2;
|
|
} else {
|
|
D += 1;
|
|
}
|
|
a += 1;
|
|
k -= 1;
|
|
}
|
|
}
|
|
|
|
size_t
|
|
MLASCALL
|
|
MlasSgemmKernelPOWER10(
|
|
const float* A,
|
|
const float* B,
|
|
float* C,
|
|
size_t CountK,
|
|
size_t CountM,
|
|
size_t CountN,
|
|
size_t lda,
|
|
size_t ldc,
|
|
float alpha,
|
|
bool ZeroMode
|
|
)
|
|
/*++
|
|
|
|
Routine Description:
|
|
|
|
This routine is an inner kernel to compute matrix multiplication for a
|
|
set of rows.
|
|
|
|
Arguments:
|
|
|
|
A - Supplies the address of matrix A.
|
|
|
|
B - Supplies the address of matrix B. The matrix data has been packed using
|
|
MlasSgemmCopyPackB or MlasSgemmTransposePackB.
|
|
|
|
C - Supplies the address of matrix C.
|
|
|
|
CountK - Supplies the number of columns from matrix A and the number of rows
|
|
from matrix B to iterate over.
|
|
|
|
CountM - Supplies the maximum number of rows that can be processed for
|
|
matrix A and matrix C. The actual number of rows handled for this
|
|
invocation depends on the kernel implementation.
|
|
|
|
CountN - Supplies the number of columns from matrix B and matrix C to
|
|
iterate over.
|
|
|
|
lda - Supplies the first dimension of matrix A.
|
|
|
|
ldc - Supplies the first dimension of matrix C.
|
|
|
|
alpha - Supplies the scalar multiplier (see SGEMM definition).
|
|
|
|
ZeroMode - Supplies true if the output matrix must be zero initialized,
|
|
else false if the output matrix is accumulated into.
|
|
|
|
Return Value:
|
|
|
|
Returns the number of rows handled.
|
|
|
|
--*/
|
|
{
|
|
size_t RowsHandled;
|
|
size_t index = CountK * 2;
|
|
|
|
MLAS_FLOAT32X4* PackA =
|
|
reinterpret_cast<MLAS_FLOAT32X4*>(alloca(sizeof(MLAS_FLOAT32X4) * index));
|
|
MLAS_FLOAT32X4 AlphaBroadcast = MlasBroadcastFloat32x4(alpha);
|
|
|
|
if (CountM >= 8) {
|
|
#ifdef _AIX
|
|
MlasSgemmPackA<8>(PackA, A, lda, CountK);
|
|
#else
|
|
if (CountK >= 16 && !(CountK % 16)) {
|
|
PackAKernelPOWER10(PackA, A, lda, CountK, 8);
|
|
} else {
|
|
MlasSgemmPackA<8>(PackA, A, lda, CountK);
|
|
}
|
|
#endif
|
|
|
|
RowsHandled = MlasSgemmMMAProcessCount<4>(PackA, B, C, 8, CountK, CountN, ldc, AlphaBroadcast, ZeroMode);
|
|
} else if (CountM >= 4) {
|
|
memset(PackA + CountK, 0, sizeof(MLAS_FLOAT32X4) * CountK);
|
|
#ifdef _AIX
|
|
MlasSgemmPackA<4>(PackA, A, lda, CountK);
|
|
#else
|
|
if (CountK >= 16 && !(CountK % 16)) {
|
|
PackAKernelPOWER10(PackA, A, lda, CountK, 4);
|
|
} else {
|
|
MlasSgemmPackA<4>(PackA, A, lda, CountK);
|
|
}
|
|
#endif
|
|
RowsHandled = MlasSgemmMMAProcessCount<4>(PackA, B, C, 4, CountK, CountN, ldc, AlphaBroadcast, ZeroMode);
|
|
} else if (CountM >= 2) {
|
|
RowsHandled = MlasSgemmProcessCount<2>(A, B, C, CountK, CountN, lda, ldc, AlphaBroadcast, ZeroMode);
|
|
} else {
|
|
RowsHandled = MlasSgemmProcessCount<1>(A, B, C, CountK, CountN, lda, ldc, AlphaBroadcast, ZeroMode);
|
|
}
|
|
return RowsHandled;
|
|
}
|