Add Fastor library

This commit is contained in:
Bassem Girgis
2025-03-22 01:17:52 -05:00
parent 5546e086f6
commit 4dd5939693
132 changed files with 55086 additions and 0 deletions

View File

@@ -0,0 +1,448 @@
#ifndef TRANSPOSE_H
#define TRANSPOSE_H
#include "Fastor/config/config.h"
#include "Fastor/backend/transpose/transpose_kernels.h"
#include "Fastor/simd_vector/extintrin.h"
#include "Fastor/simd_vector/SIMDVector.h"
namespace Fastor {
// Forward declare
namespace internal {
template<typename T, size_t M, size_t N>
FASTOR_INLINE void _transpose_dispatch(const T * FASTOR_RESTRICT a, T * FASTOR_RESTRICT out);
} // internal
//----------------------------------------------------------------------------------------------------------//
#ifdef FASTOR_AVX_IMPL
template<typename T, size_t M, size_t N>
FASTOR_INLINE void _transpose(const T * FASTOR_RESTRICT a, T * FASTOR_RESTRICT out) {
using V = SIMDVector<T,DEFAULT_ABI>;
// Block sizes of 8x8 i.e. numSIMDRows=1
// numSIMDCols=1 and innerBlock=outerBlock=1
// give a much greater speed up, but causes
// significant slow-down for issue #42
#ifndef FASTOR_TRANS_OUTER_BLOCK_SIZE
constexpr size_t numSIMDRows = 1UL;
#else
constexpr size_t numSIMDRows = FASTOR_TRANS_OUTER_BLOCK_SIZE;
#endif
#ifndef FASTOR_TRANS_INNER_BLOCK_SIZE
constexpr size_t numSIMDCols = 1UL;
#else
constexpr size_t numSIMDCols = FASTOR_TRANS_INNER_BLOCK_SIZE;
#endif
constexpr size_t innerBlock = V::Size * numSIMDCols;
constexpr size_t outerBlock = V::Size * numSIMDRows;
FASTOR_ARCH_ALIGN T pack_a[outerBlock*innerBlock];
FASTOR_ARCH_ALIGN T pack_out[outerBlock*innerBlock];
constexpr size_t M0 = M / innerBlock * innerBlock;
constexpr size_t N0 = N / outerBlock * outerBlock;
V _vec;
// For row-major matrices we go over N
// and then M to get contiguous writes
size_t j=0;
for (; j<N0; j+=outerBlock) {
size_t i=0;
for (; i< M0; i+=innerBlock) {
// Pack A
for (size_t ii=0; ii<innerBlock; ++ii) {
_vec.load(&a[(i+ii)*N+(j)],false);
_vec.store(&pack_a[ii*outerBlock]);
}
// Perform transpose on pack_a and get the result
// on pack_out
internal::_transpose_dispatch<T,innerBlock,outerBlock>(pack_a,pack_out);
// Unpack pack_out to out
for (size_t jj=0; jj<outerBlock; ++jj) {
_vec.load(&pack_out[jj*innerBlock]);
_vec.store(&out[(j+jj)*M+(i)],false);
}
}
// Remainer M - M0 columns (of c)
for (; i< M; ++i) {
for (size_t jj=0; jj<outerBlock; ++jj) {
out[(j+jj)*M+(i)] = a[i*N+j+jj];
}
}
}
// Remainder N - N0 rows (of c)
for (; j<N; ++j) {
for (size_t i=0; i< M; ++i) {
out[(j)*M+(i)] = a[i*N+j];
}
}
}
#else
template<typename T, size_t M, size_t N>
FASTOR_INLINE void _transpose(const T * FASTOR_RESTRICT a, T * FASTOR_RESTRICT out) {
for (size_t j=0; j<N; ++j)
for (size_t i=0; i< M; ++i)
out[j*M+i] = a[i*N+j];
}
#endif
//----------------------------------------------------------------------------------------------------------//
// Specialisations - float
//----------------------------------------------------------------------------------------------------------//
#ifdef FASTOR_SSE2_IMPL
template<>
FASTOR_INLINE void _transpose<float,2,2>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
__m128 a_reg = _mm_loadu_ps(a);
_mm_storeu_ps(out,_mm_shuffle_ps(a_reg,a_reg,_MM_SHUFFLE(3,1,2,0)));
}
template<>
FASTOR_INLINE void _transpose<float,3,3>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
#ifndef FASTOR_AVX2_IMPL
// 5 OPS
__m128 row0 = _mm_loadu_ps(a);
__m128 row1 = _mm_loadu_ps(a+3);
__m128 row2 = _mm_loadu_ps(a+6);
__m128 T0 = _mm_unpacklo_ps(row0,row1);
__m128 T1 = _mm_unpackhi_ps(row0,row1);
row0 = _mm_movelh_ps ( T0,row2 );
row1 = _mm_shuffle_ps( T0,row2, _MM_SHUFFLE(3,1,3,2) );
row2 = _mm_shuffle_ps( T1,row2, _MM_SHUFFLE(3,2,1,0) );
_mm_storeu_ps(out,row0);
_mm_storeu_ps(out+3,row1);
_mm_storeu_ps(out+6,row2); // out of range for out[9]
#else
// 3 OPS
// gcc/clang emit vpermsps tht operate on (%rsp)
// less pressure on shuffle port perhaps
__m256 trans07 = _mm256_loadu_ps(a);
const __m256i trans_mask = _mm256_setr_epi32(
0,3,6,
1,4,7,
2,5);
// does not shuffle across 256 lanes, only 128 lanes
// __m256 _res = _mm256_permutevar_ps(trans07, trans_mask);
// this one shuffles correctly
__m256 _res = _mm256_permutevar8x32_ps(trans07, trans_mask);
_mm256_storeu_ps(out,_res);
// out[8] = a[8];
_mm_store_ss(out+8,_mm_load_ss(a+8));
#endif
}
template<>
FASTOR_INLINE void _transpose<float,4,4>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
#ifdef FASTOR_AVX512F_IMPL
__m512 amm = _mm512_loadu_ps(a);
__m512i idx = _mm512_setr_epi32( 0, 4, 8, 12,
1, 5, 9, 13,
2, 6, 10, 14,
3, 7, 11, 15);
__m512 omm = _mm512_permutexvar_ps(idx, amm);
_mm512_storeu_ps(out, omm);
#else
__m128 row1 = _mm_loadu_ps(a);
__m128 row2 = _mm_loadu_ps(a+4);
__m128 row3 = _mm_loadu_ps(a+8);
__m128 row4 = _mm_loadu_ps(a+12);
_MM_TRANSPOSE4_PS(row1, row2, row3, row4);
_mm_storeu_ps(out , row1);
_mm_storeu_ps(out+4 , row2);
_mm_storeu_ps(out+8 , row3);
_mm_storeu_ps(out+12, row4);
#endif
}
#endif
#ifdef FASTOR_AVX_IMPL
template<>
FASTOR_INLINE void _transpose<float,8,8>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
__m256 row1 = _mm256_loadu_ps(a);
__m256 row2 = _mm256_loadu_ps(a+8);
__m256 row3 = _mm256_loadu_ps(a+16);
__m256 row4 = _mm256_loadu_ps(a+24);
__m256 row5 = _mm256_loadu_ps(a+32);
__m256 row6 = _mm256_loadu_ps(a+40);
__m256 row7 = _mm256_loadu_ps(a+48);
__m256 row8 = _mm256_loadu_ps(a+56);
internal::_MM_TRANSPOSE8_PS(row1, row2, row3, row4, row5, row6, row7, row8);
_mm256_storeu_ps(out, row1);
_mm256_storeu_ps(out+8, row2);
_mm256_storeu_ps(out+16, row3);
_mm256_storeu_ps(out+24, row4);
_mm256_storeu_ps(out+32, row5);
_mm256_storeu_ps(out+40, row6);
_mm256_storeu_ps(out+48, row7);
_mm256_storeu_ps(out+56, row8);
}
#endif
#if defined(FASTOR_AVX512F_IMPL) && defined(FASTOR_AVX512DQ_IMPL)
template<>
FASTOR_INLINE void _transpose<float,16,16>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
internal::_MM_TRANSPOSE16_PS(a,out);
}
#endif
//----------------------------------------------------------------------------------------------------------//
// Specialisations - double
//----------------------------------------------------------------------------------------------------------//
#ifdef FASTOR_SSE2_IMPL
template<>
FASTOR_INLINE void _transpose<double,2,2>(const double* FASTOR_RESTRICT a, double* FASTOR_RESTRICT out) {
/*-------------------------------------------------------*/
// 2 OPS
__m128d row0 = _mm_loadu_pd(a);
__m128d row1 = _mm_loadu_pd(a+2);
__m128d tmp = row0;
row0 = _mm_shuffle_pd(row0,row1,0x0);
row1 = _mm_shuffle_pd(tmp ,row1,0x3);
_mm_storeu_pd(out ,row0);
_mm_storeu_pd(out+2,row1);
/*-------------------------------------------------------*/
/*-------------------------------------------------------*/
// // AVX VERSION
// // IVY 4 OPS / HW 8 OPS
// __m256d a1 = _mm256_loadu_pd(a);
// __m128d a2 = _mm256_castpd256_pd128(a1);
// __m128d a3 = _mm256_extractf128_pd(a1,0x1);
// __m128d a4 = _mm_shuffle_pd(a2,a3,0x0);
// a3 = _mm_shuffle_pd(a2,a3,0x3);
// a1 = _mm256_castpd128_pd256(a4);
// a1 = _mm256_insertf128_pd(a1,a3,0x1);
// _mm256_storeu_pd(out,a1);
/*-------------------------------------------------------*/
}
#endif
#ifdef FASTOR_SSE2_IMPL
template<>
FASTOR_INLINE void _transpose<double,3,3>(const double* FASTOR_RESTRICT a, double* FASTOR_RESTRICT out) {
// AVX512 is the fastest & AVX version despite more instructions is faster than SSE
#if defined(FASTOR_AVX512F_IMPL)
// AVX512
/*-------------------------------------------------------*/
__m512d a07 = _mm512_loadu_pd(a);
const __m512i trans_mask = _mm512_setr_epi64(
0,3,6,
1,4,7,
2,5);
__m512d trans07 = _mm512_permutexvar_pd(trans_mask, a07);
_mm512_storeu_pd(out,trans07);
_mm_store_sd(out+8,_mm_load_sd(a+8));
/*-------------------------------------------------------*/
#elif defined(FASTOR_AVX_IMPL)
// AVX
/*-------------------------------------------------------*/
__m256d row1 = _mm256_loadu_pd(a);
__m256d row2 = _mm256_loadu_pd(a+4);
__m128d a11 = _mm256_castpd256_pd128(row1);
__m128d a12 = _mm256_extractf128_pd(row1,0x1);
__m128d a21 = _mm256_castpd256_pd128(row2);
__m128d a22 = _mm256_extractf128_pd(row2,0x1);
row1 = _mm256_castpd128_pd256(_mm_shuffle_pd(a11,a12,0x2));
row1 = _mm256_insertf128_pd(row1,_mm_shuffle_pd(a22,a11,0x2),0x1);
row2 = _mm256_castpd128_pd256(_mm_shuffle_pd(a21,a22,0x2));
row2 = _mm256_insertf128_pd(row2,_mm_shuffle_pd(a12,a21,0x2),0x1);
_mm256_storeu_pd(out,row1);
_mm256_storeu_pd(out+4,row2);
_mm_store_sd(out+8,_mm_load_sd(a+8));
/*-------------------------------------------------------*/
#else
// SSE
/*-------------------------------------------------------*/
__m128d a11 = _mm_loadu_pd(a);
__m128d a12 = _mm_loadu_pd(a+2);
__m128d a21 = _mm_loadu_pd(a+4);
__m128d a22 = _mm_loadu_pd(a+6);
_mm_storeu_pd(out ,_mm_shuffle_pd(a11,a12,0x2));
_mm_storeu_pd(out+2,_mm_shuffle_pd(a22,a11,0x2));
_mm_storeu_pd(out+4,_mm_shuffle_pd(a21,a22,0x2));
_mm_storeu_pd(out+6,_mm_shuffle_pd(a12,a21,0x2));
_mm_store_sd (out+8,_mm_load_sd(a+8));
/*-------------------------------------------------------*/
#endif
}
#endif
#ifdef FASTOR_AVX_IMPL
template<>
FASTOR_INLINE void _transpose<double,4,4>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
#ifdef FASTOR_AVX512F_IMPL
__m512d amm0 = _mm512_loadu_pd(a);
__m512d amm1 = _mm512_loadu_pd(a+8);
__m512i idx0 = _mm512_setr_epi64(0, 4, 8, 12, 1, 5, 9, 13);
__m512i idx1 = _mm512_setr_epi64(2, 6, 10, 14, 3, 7, 11, 15);
__m512d omm0 = _mm512_permutex2var_pd(amm0, idx0, amm1);
__m512d omm1 = _mm512_permutex2var_pd(amm0, idx1, amm1);
_mm512_storeu_pd(out , omm0);
_mm512_storeu_pd(out+8, omm1);
#else
__m256d row1 = _mm256_loadu_pd(a);
__m256d row2 = _mm256_loadu_pd(a+4);
__m256d row3 = _mm256_loadu_pd(a+8);
__m256d row4 = _mm256_loadu_pd(a+12);
internal::_MM_TRANSPOSE4_PD(row1, row2, row3, row4);
_mm256_storeu_pd(out, row1);
_mm256_storeu_pd(out+4, row2);
_mm256_storeu_pd(out+8, row3);
_mm256_storeu_pd(out+12, row4);
#endif
}
#endif
#if defined(FASTOR_AVX512F_IMPL)
template<>
FASTOR_INLINE void _transpose<double,8,8>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
__m512d row0 = _mm512_loadu_pd(a);
__m512d row1 = _mm512_loadu_pd(a+8);
__m512d row2 = _mm512_loadu_pd(a+16);
__m512d row3 = _mm512_loadu_pd(a+24);
__m512d row4 = _mm512_loadu_pd(a+32);
__m512d row5 = _mm512_loadu_pd(a+40);
__m512d row6 = _mm512_loadu_pd(a+48);
__m512d row7 = _mm512_loadu_pd(a+56);
internal::_MM_TRANSPOSE8_PD(row0,row1,row2,row3,row4,row5,row6,row7);
_mm512_storeu_pd(out , row0);
_mm512_storeu_pd(out+8 , row1);
_mm512_storeu_pd(out+16, row2);
_mm512_storeu_pd(out+24, row3);
_mm512_storeu_pd(out+32, row4);
_mm512_storeu_pd(out+40, row5);
_mm512_storeu_pd(out+48, row6);
_mm512_storeu_pd(out+56, row7);
}
#elif defined(FASTOR_AVX_IMPL)
template<>
FASTOR_INLINE void _transpose<double,8,8>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
{
__m256d row1 = _mm256_loadu_pd(&a[0]);
__m256d row2 = _mm256_loadu_pd(&a[8]);
__m256d row3 = _mm256_loadu_pd(&a[16]);
__m256d row4 = _mm256_loadu_pd(&a[24]);
internal::_MM_TRANSPOSE4_PD(row1, row2, row3, row4);
_mm256_storeu_pd(&out[0], row1);
_mm256_storeu_pd(&out[8], row2);
_mm256_storeu_pd(&out[16], row3);
_mm256_storeu_pd(&out[24], row4);
}
{
__m256d row1 = _mm256_loadu_pd(&a[32]);
__m256d row2 = _mm256_loadu_pd(&a[40]);
__m256d row3 = _mm256_loadu_pd(&a[48]);
__m256d row4 = _mm256_loadu_pd(&a[56]);
internal::_MM_TRANSPOSE4_PD(row1, row2, row3, row4);
_mm256_storeu_pd(&out[4], row1);
_mm256_storeu_pd(&out[12], row2);
_mm256_storeu_pd(&out[20], row3);
_mm256_storeu_pd(&out[28], row4);
}
{
__m256d row1 = _mm256_loadu_pd(&a[4]);
__m256d row2 = _mm256_loadu_pd(&a[12]);
__m256d row3 = _mm256_loadu_pd(&a[20]);
__m256d row4 = _mm256_loadu_pd(&a[28]);
internal::_MM_TRANSPOSE4_PD(row1, row2, row3, row4);
_mm256_storeu_pd(&out[32], row1);
_mm256_storeu_pd(&out[40], row2);
_mm256_storeu_pd(&out[48], row3);
_mm256_storeu_pd(&out[56], row4);
}
{
__m256d row1 = _mm256_loadu_pd(&a[36]);
__m256d row2 = _mm256_loadu_pd(&a[44]);
__m256d row3 = _mm256_loadu_pd(&a[52]);
__m256d row4 = _mm256_loadu_pd(&a[60]);
internal::_MM_TRANSPOSE4_PD(row1, row2, row3, row4);
_mm256_storeu_pd(&out[36], row1);
_mm256_storeu_pd(&out[44], row2);
_mm256_storeu_pd(&out[52], row3);
_mm256_storeu_pd(&out[60], row4);
}
}
#endif
//----------------------------------------------------------------------------------------------------------//
//----------------------------------------------------------------------------------------------------------//
namespace internal {
// To get around compilers recusive inlining depth issue
template<typename T, size_t M, size_t N>
FASTOR_INLINE void _transpose_dispatch(const T * FASTOR_RESTRICT a, T * FASTOR_RESTRICT out) {
for (size_t j=0; j<N; ++j)
for (size_t i=0; i< M; ++i)
out[j*M+i] = a[i*N+j];
}
template<>
FASTOR_INLINE void _transpose_dispatch<float,2,2>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
_transpose<float,2,2>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<float,3,3>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
_transpose<float,3,3>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<float,4,4>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
_transpose<float,4,4>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<float,8,8>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
_transpose<float,8,8>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<float,16,16>(const float * FASTOR_RESTRICT a, float * FASTOR_RESTRICT out) {
_transpose<float,16,16>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<double,2,2>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
_transpose<double,2,2>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<double,3,3>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
_transpose<double,3,3>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<double,4,4>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
_transpose<double,4,4>(a,out);
}
template<>
FASTOR_INLINE void _transpose_dispatch<double,8,8>(const double * FASTOR_RESTRICT a, double * FASTOR_RESTRICT out) {
_transpose<double,8,8>(a,out);
}
} // internal
//----------------------------------------------------------------------------------------------------------//
}
#endif // TRANSPOSE_H

View File

@@ -0,0 +1,233 @@
#ifndef TRANSPOSE_KERNELS_H
#define TRANSPOSE_KERNELS_H
#include "Fastor/config/config.h"
#include "Fastor/simd_vector/extintrin.h"
namespace Fastor {
namespace internal {
#ifdef FASTOR_SSE_IMPL
// 4x4 PS - defined
// _MM_TRANSPOSE4_PS
#endif
#ifdef FASTOR_AVX_IMPL
// 4x4 PD
FASTOR_INLINE void _MM_TRANSPOSE4_PD(__m256d &row0, __m256d &row1, __m256d &row2, __m256d &row3)
{
__m256d tmp3, tmp2, tmp1, tmp0;
tmp0 = _mm256_shuffle_pd((row0),(row1), 0x0);
tmp2 = _mm256_shuffle_pd((row0),(row1), 0xF);
tmp1 = _mm256_shuffle_pd((row2),(row3), 0x0);
tmp3 = _mm256_shuffle_pd((row2),(row3), 0xF);
row0 = _mm256_permute2f128_pd(tmp0, tmp1, 0x20);
row1 = _mm256_permute2f128_pd(tmp2, tmp3, 0x20);
row2 = _mm256_permute2f128_pd(tmp0, tmp1, 0x31);
row3 = _mm256_permute2f128_pd(tmp2, tmp3, 0x31);
}
// 8x8 PS
FASTOR_INLINE void _MM_TRANSPOSE8_PS(__m256 &row0, __m256 &row1, __m256 &row2, __m256 &row3,
__m256 &row4, __m256 &row5, __m256 &row6, __m256 &row7)
{
__m256 __t0, __t1, __t2, __t3, __t4, __t5, __t6, __t7;
__m256 __tt0, __tt1, __tt2, __tt3, __tt4, __tt5, __tt6, __tt7;
__t0 = _mm256_unpacklo_ps(row0, row1);
__t1 = _mm256_unpackhi_ps(row0, row1);
__t2 = _mm256_unpacklo_ps(row2, row3);
__t3 = _mm256_unpackhi_ps(row2, row3);
__t4 = _mm256_unpacklo_ps(row4, row5);
__t5 = _mm256_unpackhi_ps(row4, row5);
__t6 = _mm256_unpacklo_ps(row6, row7);
__t7 = _mm256_unpackhi_ps(row6, row7);
__tt0 = _mm256_shuffle_ps(__t0,__t2,_MM_SHUFFLE(1,0,1,0));
__tt1 = _mm256_shuffle_ps(__t0,__t2,_MM_SHUFFLE(3,2,3,2));
__tt2 = _mm256_shuffle_ps(__t1,__t3,_MM_SHUFFLE(1,0,1,0));
__tt3 = _mm256_shuffle_ps(__t1,__t3,_MM_SHUFFLE(3,2,3,2));
__tt4 = _mm256_shuffle_ps(__t4,__t6,_MM_SHUFFLE(1,0,1,0));
__tt5 = _mm256_shuffle_ps(__t4,__t6,_MM_SHUFFLE(3,2,3,2));
__tt6 = _mm256_shuffle_ps(__t5,__t7,_MM_SHUFFLE(1,0,1,0));
__tt7 = _mm256_shuffle_ps(__t5,__t7,_MM_SHUFFLE(3,2,3,2));
row0 = _mm256_permute2f128_ps(__tt0, __tt4, 0x20);
row1 = _mm256_permute2f128_ps(__tt1, __tt5, 0x20);
row2 = _mm256_permute2f128_ps(__tt2, __tt6, 0x20);
row3 = _mm256_permute2f128_ps(__tt3, __tt7, 0x20);
row4 = _mm256_permute2f128_ps(__tt0, __tt4, 0x31);
row5 = _mm256_permute2f128_ps(__tt1, __tt5, 0x31);
row6 = _mm256_permute2f128_ps(__tt2, __tt6, 0x31);
row7 = _mm256_permute2f128_ps(__tt3, __tt7, 0x31);
}
#endif
#ifdef FASTOR_AVX512F_IMPL
// 8x8 PD
inline void _MM_TRANSPOSE8_PD(__m512d &row0, __m512d &row1, __m512d &row2, __m512d &row3,
__m512d &row4, __m512d &row5, __m512d &row6, __m512d &row7)
{
__m512d __t0, __t1, __t2, __t3, __t4, __t5, __t6, __t7;
__m512d __tt0, __tt1, __tt2, __tt3, __tt4, __tt5, __tt6, __tt7;
FASTOR_ARCH_ALIGN constexpr int64_t idx1[8] = {0, 8 , 1 , 9 , 4 , 12, 5 , 13};
FASTOR_ARCH_ALIGN constexpr int64_t idx2[8] = {2, 10, 3 , 11, 6 , 14, 7 , 15};
FASTOR_ARCH_ALIGN constexpr int64_t idx3[8] = {0, 1 , 8 , 9 , 4 , 5 , 12, 13};
FASTOR_ARCH_ALIGN constexpr int64_t idx4[8] = {2, 3 , 10, 11, 6 , 7 , 14, 15};
FASTOR_ARCH_ALIGN constexpr int64_t idx5[8] = {4, 5 , 6 , 7 , 12, 13, 14, 15};
__m512i vidx1 = _mm512_load_epi64(idx1);
__m512i vidx2 = _mm512_load_epi64(idx2);
__m512i vidx3 = _mm512_load_epi64(idx3);
__m512i vidx4 = _mm512_load_epi64(idx4);
__m512i vidx5 = _mm512_load_epi64(idx5);
__t0 = _mm512_permutex2var_pd(row0, vidx1, row1);
__t1 = _mm512_permutex2var_pd(row0, vidx2, row1);
__t2 = _mm512_permutex2var_pd(row2, vidx1, row3);
__t3 = _mm512_permutex2var_pd(row2, vidx2, row3);
__t4 = _mm512_permutex2var_pd(row4, vidx1, row5);
__t5 = _mm512_permutex2var_pd(row4, vidx2, row5);
__t6 = _mm512_permutex2var_pd(row6, vidx1, row7);
__t7 = _mm512_permutex2var_pd(row6, vidx2, row7);
__tt0 = _mm512_permutex2var_pd(__t0, vidx3, __t2);
__tt1 = _mm512_permutex2var_pd(__t0, vidx4, __t2);
__tt2 = _mm512_permutex2var_pd(__t1, vidx3, __t3);
__tt3 = _mm512_permutex2var_pd(__t1, vidx4, __t3);
__tt4 = _mm512_permutex2var_pd(__t4, vidx3, __t6);
__tt5 = _mm512_permutex2var_pd(__t4, vidx4, __t6);
__tt6 = _mm512_permutex2var_pd(__t5, vidx3, __t7);
__tt7 = _mm512_permutex2var_pd(__t5, vidx4, __t7);
row0 = _mm512_insertf64x4(__tt0,_mm512_castpd512_pd256(__tt4),0x1);
row1 = _mm512_insertf64x4(__tt1,_mm512_castpd512_pd256(__tt5),0x1);
row2 = _mm512_insertf64x4(__tt2,_mm512_castpd512_pd256(__tt6),0x1);
row3 = _mm512_insertf64x4(__tt3,_mm512_castpd512_pd256(__tt7),0x1);
row4 = _mm512_permutex2var_pd(__tt0, vidx5, __tt4);
row5 = _mm512_permutex2var_pd(__tt1, vidx5, __tt5);
row6 = _mm512_permutex2var_pd(__tt2, vidx5, __tt6);
row7 = _mm512_permutex2var_pd(__tt3, vidx5, __tt7);
}
#endif
#if defined(FASTOR_AVX512F_IMPL) && defined(FASTOR_AVX512DQ_IMPL)
// 16x16
FASTOR_INLINE void _MM_TRANSPOSE16_PS(const float * FASTOR_RESTRICT mat, float * FASTOR_RESTRICT matT)
{
__m512 t0, t1, t2, t3, t4, t5, t6, t7, t8, t9, ta, tb, tc, td, te, tf;
__m512 r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf;
int mask;
FASTOR_ARCH_ALIGN constexpr int64_t idx1[8] = {2, 3, 0, 1, 6, 7, 4, 5};
FASTOR_ARCH_ALIGN constexpr int64_t idx2[8] = {1, 0, 3, 2, 5, 4, 7, 6};
FASTOR_ARCH_ALIGN constexpr int32_t idx3[16] = {1, 0, 3, 2, 5 ,4 ,7 ,6 ,9 ,8 , 11, 10, 13, 12 ,15, 14};
__m512i vidx1 = _mm512_load_epi64(idx1);
__m512i vidx2 = _mm512_load_epi64(idx2);
__m512i vidx3 = _mm512_load_epi32(idx3);
t0 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 0*16+0])), _mm256_loadu_ps(&mat[ 8*16+0]), 1);
t1 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 1*16+0])), _mm256_loadu_ps(&mat[ 9*16+0]), 1);
t2 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 2*16+0])), _mm256_loadu_ps(&mat[10*16+0]), 1);
t3 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 3*16+0])), _mm256_loadu_ps(&mat[11*16+0]), 1);
t4 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 4*16+0])), _mm256_loadu_ps(&mat[12*16+0]), 1);
t5 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 5*16+0])), _mm256_loadu_ps(&mat[13*16+0]), 1);
t6 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 6*16+0])), _mm256_loadu_ps(&mat[14*16+0]), 1);
t7 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 7*16+0])), _mm256_loadu_ps(&mat[15*16+0]), 1);
t8 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 0*16+8])), _mm256_loadu_ps(&mat[ 8*16+8]), 1);
t9 = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 1*16+8])), _mm256_loadu_ps(&mat[ 9*16+8]), 1);
ta = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 2*16+8])), _mm256_loadu_ps(&mat[10*16+8]), 1);
tb = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 3*16+8])), _mm256_loadu_ps(&mat[11*16+8]), 1);
tc = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 4*16+8])), _mm256_loadu_ps(&mat[12*16+8]), 1);
td = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 5*16+8])), _mm256_loadu_ps(&mat[13*16+8]), 1);
te = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 6*16+8])), _mm256_loadu_ps(&mat[14*16+8]), 1);
tf = _mm512_insertf32x8(_mm512_castps256_ps512(_mm256_loadu_ps(&mat[ 7*16+8])), _mm256_loadu_ps(&mat[15*16+8]), 1);
mask= 0xcc;
r0 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t0), (__mmask8)mask, vidx1, _mm512_castps_pd(t4)));
r1 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t1), (__mmask8)mask, vidx1, _mm512_castps_pd(t5)));
r2 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t2), (__mmask8)mask, vidx1, _mm512_castps_pd(t6)));
r3 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t3), (__mmask8)mask, vidx1, _mm512_castps_pd(t7)));
r8 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t8), (__mmask8)mask, vidx1, _mm512_castps_pd(tc)));
r9 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t9), (__mmask8)mask, vidx1, _mm512_castps_pd(td)));
ra = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(ta), (__mmask8)mask, vidx1, _mm512_castps_pd(te)));
rb = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(tb), (__mmask8)mask, vidx1, _mm512_castps_pd(tf)));
mask= 0x33;
r4 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t4), (__mmask8)mask, vidx1, _mm512_castps_pd(t0)));
r5 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t5), (__mmask8)mask, vidx1, _mm512_castps_pd(t1)));
r6 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t6), (__mmask8)mask, vidx1, _mm512_castps_pd(t2)));
r7 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(t7), (__mmask8)mask, vidx1, _mm512_castps_pd(t3)));
rc = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(tc), (__mmask8)mask, vidx1, _mm512_castps_pd(t8)));
rd = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(td), (__mmask8)mask, vidx1, _mm512_castps_pd(t9)));
re = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(te), (__mmask8)mask, vidx1, _mm512_castps_pd(ta)));
rf = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(tf), (__mmask8)mask, vidx1, _mm512_castps_pd(tb)));
mask = 0xaa;
t0 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r0), (__mmask8)mask, vidx2, _mm512_castps_pd(r2)));
t1 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r1), (__mmask8)mask, vidx2, _mm512_castps_pd(r3)));
t4 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r4), (__mmask8)mask, vidx2, _mm512_castps_pd(r6)));
t5 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r5), (__mmask8)mask, vidx2, _mm512_castps_pd(r7)));
t8 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r8), (__mmask8)mask, vidx2, _mm512_castps_pd(ra)));
t9 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r9), (__mmask8)mask, vidx2, _mm512_castps_pd(rb)));
tc = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(rc), (__mmask8)mask, vidx2, _mm512_castps_pd(re)));
td = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(rd), (__mmask8)mask, vidx2, _mm512_castps_pd(rf)));
mask = 0x55;
t2 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r2), (__mmask8)mask, vidx2, _mm512_castps_pd(r0)));
t3 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r3), (__mmask8)mask, vidx2, _mm512_castps_pd(r1)));
t6 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r6), (__mmask8)mask, vidx2, _mm512_castps_pd(r4)));
t7 = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(r7), (__mmask8)mask, vidx2, _mm512_castps_pd(r5)));
ta = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(ra), (__mmask8)mask, vidx2, _mm512_castps_pd(r8)));
tb = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(rb), (__mmask8)mask, vidx2, _mm512_castps_pd(r9)));
te = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(re), (__mmask8)mask, vidx2, _mm512_castps_pd(rc)));
tf = _mm512_castpd_ps(_mm512_mask_permutexvar_pd(_mm512_castps_pd(rf), (__mmask8)mask, vidx2, _mm512_castps_pd(rd)));
mask = 0xaaaa;
r0 = _mm512_mask_permutexvar_ps(t0, (__mmask16)mask, vidx3, t1);
r2 = _mm512_mask_permutexvar_ps(t2, (__mmask16)mask, vidx3, t3);
r4 = _mm512_mask_permutexvar_ps(t4, (__mmask16)mask, vidx3, t5);
r6 = _mm512_mask_permutexvar_ps(t6, (__mmask16)mask, vidx3, t7);
r8 = _mm512_mask_permutexvar_ps(t8, (__mmask16)mask, vidx3, t9);
ra = _mm512_mask_permutexvar_ps(ta, (__mmask16)mask, vidx3, tb);
rc = _mm512_mask_permutexvar_ps(tc, (__mmask16)mask, vidx3, td);
re = _mm512_mask_permutexvar_ps(te, (__mmask16)mask, vidx3, tf);
mask = 0x5555;
r1 = _mm512_mask_permutexvar_ps(t1, (__mmask16)mask, vidx3, t0);
r3 = _mm512_mask_permutexvar_ps(t3, (__mmask16)mask, vidx3, t2);
r5 = _mm512_mask_permutexvar_ps(t5, (__mmask16)mask, vidx3, t4);
r7 = _mm512_mask_permutexvar_ps(t7, (__mmask16)mask, vidx3, t6);
r9 = _mm512_mask_permutexvar_ps(t9, (__mmask16)mask, vidx3, t8);
rb = _mm512_mask_permutexvar_ps(tb, (__mmask16)mask, vidx3, ta);
rd = _mm512_mask_permutexvar_ps(td, (__mmask16)mask, vidx3, tc);
rf = _mm512_mask_permutexvar_ps(tf, (__mmask16)mask, vidx3, te);
_mm512_storeu_ps(&matT[ 0*16], r0);
_mm512_storeu_ps(&matT[ 1*16], r1);
_mm512_storeu_ps(&matT[ 2*16], r2);
_mm512_storeu_ps(&matT[ 3*16], r3);
_mm512_storeu_ps(&matT[ 4*16], r4);
_mm512_storeu_ps(&matT[ 5*16], r5);
_mm512_storeu_ps(&matT[ 6*16], r6);
_mm512_storeu_ps(&matT[ 7*16], r7);
_mm512_storeu_ps(&matT[ 8*16], r8);
_mm512_storeu_ps(&matT[ 9*16], r9);
_mm512_storeu_ps(&matT[10*16], ra);
_mm512_storeu_ps(&matT[11*16], rb);
_mm512_storeu_ps(&matT[12*16], rc);
_mm512_storeu_ps(&matT[13*16], rd);
_mm512_storeu_ps(&matT[14*16], re);
_mm512_storeu_ps(&matT[15*16], rf);
}
#endif
} // end of namespace internal
} // end of namespace Fastor
#endif // TRANSPOSE_KERNELS_H