Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions ruy/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -692,7 +692,7 @@ ruy_cc_library(

if((CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64) AND NOT MSVC)
set(ruy_7_mavx512bw_mavx512cd_mavx512dq_mavx512f_mavx512vl_arch_AVX512 ";-mavx512f;-mavx512vl;-mavx512cd;-mavx512bw;-mavx512dq")
elseif(MSVC)
elseif(MSVC AND (CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64)) # ARM64 support: /arch:AVX512 is x86-only
set(ruy_7_mavx512bw_mavx512cd_mavx512dq_mavx512f_mavx512vl_arch_AVX512 "/arch:AVX512")
else()
set(ruy_7_mavx512bw_mavx512cd_mavx512dq_mavx512f_mavx512vl_arch_AVX512 "")
Expand Down Expand Up @@ -764,7 +764,7 @@ ruy_cc_library(

if((CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64) AND NOT MSVC)
set(ruy_8_mavx2_mfma_arch_AVX2 "-mavx2;-mfma")
elseif(MSVC)
elseif(MSVC AND (CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64)) # ARM64 support: /arch:AVX2 is x86-only
set(ruy_8_mavx2_mfma_arch_AVX2 "/arch:AVX2")
else()
set(ruy_8_mavx2_mfma_arch_AVX2 "")
Expand Down Expand Up @@ -836,7 +836,7 @@ ruy_cc_library(

if((CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64) AND NOT MSVC)
set(ruy_9_mavx_arch_AVX "-mavx")
elseif(MSVC)
elseif(MSVC AND (CMAKE_SYSTEM_PROCESSOR STREQUAL x86_64 OR CMAKE_SYSTEM_PROCESSOR STREQUAL amd64)) # ARM64 support: /arch:AVX is x86-only
set(ruy_9_mavx_arch_AVX "/arch:AVX")
else()
set(ruy_9_mavx_arch_AVX "")
Expand Down
181 changes: 181 additions & 0 deletions ruy/kernel_arm.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,13 @@ namespace ruy {
RUY_INHERIT_KERNEL(Path::kStandardCpp, Path::kNeon)
RUY_INHERIT_KERNEL(Path::kNeon, Path::kNeonDotprod)

// MSVC on Windows ARM64 does not support AT&T-style inline assembly.
// The asm-based kernel function declarations and their calling template
// specializations are excluded under MSVC. The RUY_INHERIT_KERNEL chain above
// means kNeon and kNeonDotprod fall back to kStandardCpp on MSVC ARM64.
// MSVC-compatible NEON intrinsic kernels are provided in kernel_arm64_msvc.cc.
#if !defined(_MSC_VER)

#if RUY_PLATFORM_NEON_64
void Kernel8bitNeon(const KernelParams8bit<4, 4>& params);
void Kernel8bitNeon1Col(const KernelParams8bit<4, 4>& params);
Expand Down Expand Up @@ -213,6 +220,180 @@ struct Kernel<Path::kNeonDotprod, float, float, float, float> {
}
};

#endif // !defined(_MSC_VER)

#if defined(_MSC_VER) && defined(_M_ARM64)

#if RUY_PLATFORM_NEON_64
void Kernel8bitNeon(const KernelParams8bit<4, 4>& params);
void Kernel8bitNeon1Col(const KernelParams8bit<4, 4>& params);
void Kernel8bitNeonA55ish(const KernelParams8bit<4, 4>& params);
void Kernel8bitNeonDotprod(const KernelParams8bit<8, 8>& params);
void Kernel8bitNeonDotprod1Col(const KernelParams8bit<8, 8>& params);
void Kernel8bitNeonDotprodA55ish(const KernelParams8bit<8, 8>& params);
void Kernel8bitNeonDotprodX1(const KernelParams8bit<8, 8>& params);
// Mixed-precision NEON kernels: i8×i16 and i16×i8, tile 4×4, depth step 8.
void Kernel8bitNeonMixedInt16Lhs(const KernelParams8bit<4, 4>& params);
void Kernel8bitNeonMixedInt16Rhs(const KernelParams8bit<4, 4>& params);

template <typename DstScalar>
struct Kernel<Path::kNeon, std::int8_t, std::int8_t, std::int32_t, DstScalar> {
static constexpr Path kPath = Path::kNeon;
using LhsLayout = FixedKernelLayout<Order::kColMajor, 16, 4>;
using RhsLayout = FixedKernelLayout<Order::kColMajor, 16, 4>;
Tuning tuning = Tuning::kAuto;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<std::int8_t>& lhs, const PMat<std::int8_t>& rhs,
const MulParams<std::int32_t, DstScalar>& mul_params, int start_row,
int start_col, int end_row, int end_col, Mat<DstScalar>* dst) const {
KernelParams8bit<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParams8bit(lhs, rhs, mul_params, start_row, start_col, end_row,
end_col, dst, &params);
if (dst->layout.cols == 1 &&
mul_params.channel_dimension() == ChannelDimension::kRow) {
Kernel8bitNeon1Col(params);
return;
}
if (__builtin_expect(tuning == Tuning::kA55ish, true)) {
Kernel8bitNeonA55ish(params);
} else {
Kernel8bitNeon(params);
}
}
};

template <typename DstScalar>
struct Kernel<Path::kNeonDotprod, std::int8_t, std::int8_t, std::int32_t,
DstScalar> {
static constexpr Path kPath = Path::kNeonDotprod;
Tuning tuning = Tuning::kAuto;
using LhsLayout = FixedKernelLayout<Order::kColMajor, 4, 8>;
using RhsLayout = FixedKernelLayout<Order::kColMajor, 4, 8>;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<std::int8_t>& lhs, const PMat<std::int8_t>& rhs,
const MulParams<std::int32_t, DstScalar>& mul_params, int start_row,
int start_col, int end_row, int end_col, Mat<DstScalar>* dst) const {
KernelParams8bit<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParams8bit(lhs, rhs, mul_params, start_row, start_col, end_row,
end_col, dst, &params);
if (dst->layout.cols == 1 &&
mul_params.channel_dimension() == ChannelDimension::kRow) {
Kernel8bitNeonDotprod1Col(params);
} else if (__builtin_expect(tuning == Tuning::kA55ish, true)) {
Kernel8bitNeonDotprodA55ish(params);
} else if (tuning == Tuning::kX1) {
Kernel8bitNeonDotprodX1(params);
} else {
Kernel8bitNeonDotprod(params);
}
}
};
#endif // RUY_PLATFORM_NEON_64

void KernelFloatNeon(const KernelParamsFloat<8, 8>& params);
void KernelFloatNeonX1(const KernelParamsFloat<8, 8>& params);
void KernelFloatNeonA55ish(const KernelParamsFloat<8, 8>& params);
void KernelFloatNeonDotprodA55ish(const KernelParamsFloat<8, 8>& params);

#if RUY_PLATFORM_NEON_64
template <>
struct Kernel<Path::kNeon, float, float, float, float> {
static constexpr Path kPath = Path::kNeon;
Tuning tuning = Tuning::kAuto;
using LhsLayout = FixedKernelLayout<Order::kRowMajor, 1, 8>;
using RhsLayout = FixedKernelLayout<Order::kRowMajor, 1, 8>;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<float>& lhs, const PMat<float>& rhs,
const MulParams<float, float>& mul_params, int start_row,
int start_col, int end_row, int end_col, Mat<float>* dst) const {
KernelParamsFloat<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParamsFloat(lhs, rhs, mul_params, start_row, start_col, end_row,
end_col, dst, &params);
if (__builtin_expect(tuning == Tuning::kA55ish, true)) {
KernelFloatNeonA55ish(params);
} else if (tuning == Tuning::kX1) {
KernelFloatNeonX1(params);
} else {
KernelFloatNeon(params);
}
}
};

template <>
struct Kernel<Path::kNeonDotprod, float, float, float, float> {
static constexpr Path kPath = Path::kNeonDotprod;
Tuning tuning = Tuning::kAuto;
using LhsLayout = FixedKernelLayout<Order::kRowMajor, 1, 8>;
using RhsLayout = FixedKernelLayout<Order::kRowMajor, 1, 8>;
using Base = Kernel<Path::kNeon, float, float, float, float>;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<float>& lhs, const PMat<float>& rhs,
const MulParams<float, float>& mul_params, int start_row,
int start_col, int end_row, int end_col, Mat<float>* dst) const {
KernelParamsFloat<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParamsFloat(lhs, rhs, mul_params, start_row, start_col, end_row,
end_col, dst, &params);
if (__builtin_expect(tuning == Tuning::kA55ish, true)) {
KernelFloatNeonDotprodA55ish(params);
} else if (tuning == Tuning::kX1) {
KernelFloatNeonX1(params);
} else {
KernelFloatNeon(params);
}
}
};
#endif // RUY_PLATFORM_NEON_64

// ---------------------------------------------------------------------------
// Mixed-precision kNeon specializations: i8×i16→i16 and i16×i8→i16.
// Layout kColMajor/8/4: 8 int16 depth values × 4 cols.
// ---------------------------------------------------------------------------
#if RUY_PLATFORM_NEON_64

// i8(LHS) × i16(RHS) → i16 output
template <>
struct Kernel<Path::kNeon, std::int8_t, std::int16_t, std::int32_t,
std::int16_t> {
static constexpr Path kPath = Path::kNeon;
using LhsLayout = FixedKernelLayout<Order::kColMajor, 8, 4>;
using RhsLayout = FixedKernelLayout<Order::kColMajor, 8, 4>;
Tuning tuning = Tuning::kAuto;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<std::int8_t>& lhs, const PMat<std::int16_t>& rhs,
const MulParams<std::int32_t, std::int16_t>& mul_params,
int start_row, int start_col, int end_row, int end_col,
Mat<std::int16_t>* dst) const {
KernelParams8bit<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParams8bit(lhs, rhs, mul_params, start_row, start_col,
end_row, end_col, dst, &params);
Kernel8bitNeonMixedInt16Rhs(params);
}
};

// i16(LHS) × i8(RHS) → i16 output
template <>
struct Kernel<Path::kNeon, std::int16_t, std::int8_t, std::int32_t,
std::int16_t> {
static constexpr Path kPath = Path::kNeon;
using LhsLayout = FixedKernelLayout<Order::kColMajor, 8, 4>;
using RhsLayout = FixedKernelLayout<Order::kColMajor, 8, 4>;
Tuning tuning = Tuning::kAuto;
explicit Kernel(Tuning tuning_) : tuning(tuning_) {}
void Run(const PMat<std::int16_t>& lhs, const PMat<std::int8_t>& rhs,
const MulParams<std::int32_t, std::int16_t>& mul_params,
int start_row, int start_col, int end_row, int end_col,
Mat<std::int16_t>* dst) const {
KernelParams8bit<LhsLayout::kCols, RhsLayout::kCols> params;
MakeKernelParams8bitMixed(lhs, rhs, mul_params, start_row, start_col,
end_row, end_col, dst, &params);
Kernel8bitNeonMixedInt16Lhs(params);
}
};

#endif // RUY_PLATFORM_NEON_64

#endif // defined(_MSC_VER) && defined(_M_ARM64)

#endif // RUY_PLATFORM_NEON && RUY_OPT(ASM)

} // namespace ruy
Expand Down
7 changes: 5 additions & 2 deletions ruy/kernel_arm32.cc
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,10 @@ limitations under the License.

namespace ruy {

#if RUY_PLATFORM_NEON_32 && RUY_OPT(ASM)
// MSVC on Windows does not support AT&T-style GCC inline assembly
// (asm volatile). ARM32 (Thumb/Thumb2) is not a target for Windows ARM64 but
// this guard keeps the file consistent and safe for any MSVC toolchain.
#if RUY_PLATFORM_NEON_32 && RUY_OPT(ASM) && !defined(_MSC_VER)

#define RUY_ASM_LABEL_STORE_UINT8 91
#define RUY_ASM_LABEL_STORE_INT8 92
Expand Down Expand Up @@ -2525,5 +2528,5 @@ void Kernel8bitNeon1Col(const KernelParams8bit<4, 2>& params) {
#undef RUY_STACK_OFFSET_LHS_COL_PTR
#undef RUY_STACK_OFFSET_RHS_COL_PTR

#endif // RUY_PLATFORM_NEON_32 && (RUY_OPT(ASM)
#endif // RUY_PLATFORM_NEON_32 && RUY_OPT(ASM) && !defined(_MSC_VER)
} // namespace ruy
Loading