diff --git a/onnxruntime/core/mlas/lib/activate.cpp b/onnxruntime/core/mlas/lib/activate.cpp index a388894bc58cd..0c5e59cdeb236 100644 --- a/onnxruntime/core/mlas/lib/activate.cpp +++ b/onnxruntime/core/mlas/lib/activate.cpp @@ -490,12 +490,19 @@ Return Value: --*/ { #if defined(MLAS_TARGET_RISCV64) && defined(MLAS_USE_RVV) - // The platform routine covers the element-wise kinds and returns false for the - // rest, which fall through to the switch below. - - if (GetMlasPlatform().ActivationRoutine != nullptr && - GetMlasPlatform().ActivationRoutine(Activation, Buffer, Bias, M, N, ldc)) { - return; + const auto activation_routine = GetMlasPlatform().ActivationRoutine; + if (activation_routine != nullptr) { +#if !defined(FORCE_GENERIC_ALGORITHMS) + // Short rows do not amortize the specialized vector setup cost. + if (N >= 32 && (Activation->ActivationKind != MlasIdentityActivation || Bias != nullptr) && + MlasFusedActivationRvv(Activation, Buffer, Bias, M, N, ldc)) { + return; + } +#endif + // Keep the existing RVV routine for short rows and unsupported kinds. + if (activation_routine(Activation, Buffer, Bias, M, N, ldc)) { + return; + } } #endif diff --git a/onnxruntime/core/mlas/lib/mlasi.h b/onnxruntime/core/mlas/lib/mlasi.h index ad9286b973d9e..66354498b7aa8 100644 --- a/onnxruntime/core/mlas/lib/mlasi.h +++ b/onnxruntime/core/mlas/lib/mlasi.h @@ -1477,6 +1477,7 @@ MlasReorderOutputNchwBlock16Avx512F( MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL MlasReduceMaximumF32Kernel; MLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL MlasReduceMinimumMaximumF32Kernel; #if defined(MLAS_TARGET_RISCV64) && defined(MLAS_USE_RVV) + MLAS_ACTIVATION_ROUTINE MlasFusedActivationRvv; MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL MlasComputeSumExpF32KernelRvv; MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL MlasReduceMaximumF32KernelRvv; MLAS_REDUCE_MINIMUM_MAXIMUM_FLOAT_KERNEL MlasReduceMinimumMaximumF32KernelRvv; diff --git a/onnxruntime/core/mlas/lib/riscv64/activation_kernel_rvv.cpp b/onnxruntime/core/mlas/lib/riscv64/activation_kernel_rvv.cpp index a53ef184e4e32..5e092c7970bf1 100644 --- a/onnxruntime/core/mlas/lib/riscv64/activation_kernel_rvv.cpp +++ b/onnxruntime/core/mlas/lib/riscv64/activation_kernel_rvv.cpp @@ -10,8 +10,8 @@ Module Name: Abstract: - RVV unary activation kernels for riscv64: erf, tanh, logistic (sigmoid), - exp, silu, gelu(erf). Wired through MLAS_PLATFORM kernel routine fields + RVV fused bias/activation and unary activation kernels for riscv64: + erf, tanh, logistic (sigmoid), exp, silu, gelu(erf). Wired through MLAS_PLATFORM fields on builds with RVV support (MLAS_USE_RVV). LMUL=m4 throughout (32 floats per vector at VLEN=256), scaling with VLEN @@ -107,6 +107,162 @@ constexpr float ERF_A5 = 1.061405429f; } // namespace +namespace { + +template +void +MlasFusedActivationKernelRvv( + const MLAS_ACTIVATION* Activation, + float* Buffer, + const float* Bias, + size_t M, + size_t N, + size_t ldc + ) +{ + float Alpha = 0.0f; + float Beta = 0.0f; + float Minimum = 0.0f; + float Maximum = 1.0f; + if constexpr (ActivationKind == MlasLeakyReluActivation) { + Alpha = Activation->Parameters.LeakyRelu.alpha; + } else if constexpr (ActivationKind == MlasClipActivation) { + Minimum = Activation->Parameters.Clip.minimum; + Maximum = Activation->Parameters.Clip.maximum; + } else if constexpr (ActivationKind == MlasHardSigmoidActivation) { + Alpha = Activation->Parameters.HardSigmoid.alpha; + Beta = Activation->Parameters.HardSigmoid.beta; + } else if constexpr (ActivationKind == MlasHardSwishActivation) { + Alpha = 1.0f / 6.0f; + Beta = 0.5f; + } + + while (M-- > 0) { + float BiasValue = 0.0f; + if constexpr (AddBias) { + BiasValue = *Bias++; + } + float* buffer = Buffer; + size_t n = N; + if constexpr (ActivationKind == MlasLeakyReluActivation) { + // The generic four-lane kernel uses > 0, but its scalar tail uses + // >= 0. Preserve that distinction for signed zero and nonfinite alpha. + n &= ~size_t(3); + } + while (n > 0) { + const size_t vl = __riscv_vsetvl_e32m4(n); + vfloat32m4_t Value = __riscv_vle32_v_f32m4(buffer, vl); + if constexpr (AddBias) { + Value = __riscv_vfadd_vf_f32m4(Value, BiasValue, vl); + } + + if constexpr (ActivationKind == MlasReluActivation) { + // Compare/select preserves NaNs and signed zero, unlike vfmax. + const vbool8_t Negative = __riscv_vmflt_vf_f32m4_b8(Value, 0.0f, vl); + Value = __riscv_vfmerge_vfm_f32m4(Value, 0.0f, Negative, vl); + } else if constexpr (ActivationKind == MlasLeakyReluActivation) { + const vfloat32m4_t Scaled = __riscv_vfmul_vf_f32m4(Value, Alpha, vl); + const vbool8_t Positive = __riscv_vmfgt_vf_f32m4_b8(Value, 0.0f, vl); + Value = __riscv_vmerge_vvm_f32m4(Scaled, Value, Positive, vl); + } else if constexpr (ActivationKind == MlasClipActivation) { + const vbool8_t Below = __riscv_vmflt_vf_f32m4_b8(Value, Minimum, vl); + Value = __riscv_vfmerge_vfm_f32m4(Value, Minimum, Below, vl); + const vbool8_t Above = __riscv_vmfgt_vf_f32m4_b8(Value, Maximum, vl); + Value = __riscv_vfmerge_vfm_f32m4(Value, Maximum, Above, vl); + } else if constexpr (ActivationKind == MlasHardSigmoidActivation || + ActivationKind == MlasHardSwishActivation) { + vfloat32m4_t Gate = __riscv_vfmul_vf_f32m4(Value, Alpha, vl); + Gate = __riscv_vfadd_vf_f32m4(Gate, Beta, vl); + const vbool8_t Above = __riscv_vmfgt_vf_f32m4_b8(Gate, Maximum, vl); + Gate = __riscv_vfmerge_vfm_f32m4(Gate, Maximum, Above, vl); + const vbool8_t Below = __riscv_vmflt_vf_f32m4_b8(Gate, Minimum, vl); + Gate = __riscv_vfmerge_vfm_f32m4(Gate, Minimum, Below, vl); + if constexpr (ActivationKind == MlasHardSwishActivation) { + Value = __riscv_vfmul_vv_f32m4(Value, Gate, vl); + } else { + Value = Gate; + } + } + + __riscv_vse32_v_f32m4(buffer, Value, vl); + buffer += vl; + n -= vl; + } + if constexpr (ActivationKind == MlasLeakyReluActivation) { + for (size_t tail = N % 4; tail > 0; --tail) { + float Value = *buffer; + if constexpr (AddBias) { + Value += BiasValue; + } + *buffer++ = (Value >= 0.0f) ? Value : Value * Alpha; + } + } + Buffer += ldc; + } +} + +template +void +MlasFusedActivationDispatchRvv( + const MLAS_ACTIVATION* Activation, + float* Buffer, + const float* Bias, + size_t M, + size_t N, + size_t ldc + ) +{ + if (Bias != nullptr) { + MlasFusedActivationKernelRvv(Activation, Buffer, Bias, M, N, ldc); + } else if constexpr (ActivationKind != MlasIdentityActivation) { + MlasFusedActivationKernelRvv(Activation, Buffer, Bias, M, N, ldc); + } +} + +} // namespace + +extern "C" +bool +MLASCALL +MlasFusedActivationRvv( + const MLAS_ACTIVATION* Activation, + float* Buffer, + const float* Bias, + size_t M, + size_t N, + size_t ldc + ) +{ + switch (Activation->ActivationKind) { + case MlasIdentityActivation: + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + case MlasReluActivation: + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + case MlasLeakyReluActivation: + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + case MlasClipActivation: + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + case MlasHardSigmoidActivation: + // Nonfinite parameters can distinguish fused from separate multiply/add. + // Leave that behavior to the generic kernel and its compilation flags. + if (!std::isfinite(Activation->Parameters.HardSigmoid.alpha) || + !std::isfinite(Activation->Parameters.HardSigmoid.beta)) { + return false; + } + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + case MlasHardSwishActivation: + MlasFusedActivationDispatchRvv(Activation, Buffer, Bias, M, N, ldc); + return true; + default: + return false; + } +} + extern "C" void MLASCALL diff --git a/onnxruntime/test/mlas/unittest/test_fused_activation.cpp b/onnxruntime/test/mlas/unittest/test_fused_activation.cpp new file mode 100644 index 0000000000000..24397604eca6c --- /dev/null +++ b/onnxruntime/test/mlas/unittest/test_fused_activation.cpp @@ -0,0 +1,251 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "test_util.h" +#include "core/mlas/lib/mlasi.h" + +#include +#include + +namespace { + +constexpr MLAS_ACTIVATION_KIND kKinds[] = { + MlasIdentityActivation, MlasReluActivation, MlasLeakyReluActivation, + MlasClipActivation, MlasHardSigmoidActivation, MlasHardSwishActivation}; + +MLAS_ACTIVATION MakeActivation(MLAS_ACTIVATION_KIND kind) { + MLAS_ACTIVATION activation{}; + activation.ActivationKind = kind; + if (kind == MlasLeakyReluActivation) { + activation.Parameters.LeakyRelu.alpha = 0.2f; + } else if (kind == MlasClipActivation) { + activation.Parameters.Clip.minimum = -0.5f; + activation.Parameters.Clip.maximum = 6.0f; + } else if (kind == MlasHardSigmoidActivation) { + activation.Parameters.HardSigmoid.alpha = 0.2f; + activation.Parameters.HardSigmoid.beta = 0.12f; + } + return activation; +} + +float Reference(const MLAS_ACTIVATION& activation, float value) { + switch (activation.ActivationKind) { + case MlasReluActivation: + return std::max(value, 0.0f); + case MlasLeakyReluActivation: + return value >= 0.0f ? value : value * activation.Parameters.LeakyRelu.alpha; + case MlasClipActivation: + return std::min(std::max(value, activation.Parameters.Clip.minimum), activation.Parameters.Clip.maximum); + case MlasHardSigmoidActivation: + return std::max(std::min(value * activation.Parameters.HardSigmoid.alpha + + activation.Parameters.HardSigmoid.beta, + 1.0f), + 0.0f); + case MlasHardSwishActivation: + return value * std::max(std::min(value * (1.0f / 6.0f) + 0.5f, 1.0f), 0.0f); + default: + return value; + } +} + +void ExpectSame(float actual, float expected) { + if (std::isnan(expected)) { + EXPECT_TRUE(std::isnan(actual)); + } else if (expected == 0.0f) { + EXPECT_EQ(actual, expected); +#if defined(MLAS_TARGET_RISCV64) + EXPECT_EQ(std::signbit(actual), std::signbit(expected)); +#endif + } else if (std::isinf(expected)) { + EXPECT_EQ(actual, expected); + } else { + EXPECT_NEAR(actual, expected, 1.0e-6f * std::max(1.0f, std::abs(expected))); + } +} + +TEST(FusedActivation, MatrixAndTail) { + constexpr size_t lengths[] = {0, 1, 2, 3, 4, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 128, 129}; + const float values[] = { + -INFINITY, -10.0f, -3.0f, -1.0f, -0.0f, 0.0f, + -std::numeric_limits::denorm_min(), std::numeric_limits::denorm_min(), + 0.25f, 1.0f, 3.0f, 10.0f, INFINITY, std::numeric_limits::quiet_NaN()}; + const float bias[] = {0.0f, -0.75f, 1.0f}; + std::array output; + std::array expected; + for (auto kind : kKinds) { + auto activation = MakeActivation(kind); + for (size_t n : lengths) { + for (size_t m : {size_t(0), size_t(1), size_t(3)}) { + for (size_t padding : {size_t(0), size_t(5)}) { + for (bool add_bias : {false, true}) { + SCOPED_TRACE(::testing::Message() << "kind=" << kind << " M=" << m << " N=" << n + << " padding=" << padding << " bias=" << add_bias); + const size_t ldc = n + padding; + output.fill(12345.0f); + expected.fill(12345.0f); + for (size_t row = 0; row < m; ++row) { + for (size_t col = 0; col < n; ++col) { + const size_t index = 1 + row * ldc + col; + output[index] = values[(row * n + col) % _countof(values)]; + const float value = add_bias ? output[index] + bias[row] : output[index]; + expected[index] = Reference(activation, value); + } + } + MlasActivation(&activation, output.data() + 1, add_bias ? bias : nullptr, m, n, ldc); + for (size_t i = 0; i < output.size(); ++i) { + ExpectSame(output[i], expected[i]); + } + } + } + } + } + } +} + +TEST(FusedActivation, GuardedTail) { + MatrixGuardBuffer guarded; + for (auto kind : kKinds) { + auto activation = MakeActivation(kind); + for (size_t n : {size_t(1), size_t(3), size_t(4), size_t(7), size_t(9), size_t(31), size_t(33), size_t(129)}) { + for (bool add_bias : {false, true}) { + SCOPED_TRACE(::testing::Message() << "kind=" << kind << " N=" << n << " bias=" << add_bias); + float* buffer = guarded.GetBuffer(n); + const float bias = 0.5f; + std::array expected; + for (size_t i = 0; i < n; ++i) { + expected[i] = Reference(activation, add_bias ? buffer[i] + bias : buffer[i]); + } + MlasActivation(&activation, buffer, add_bias ? &bias : nullptr, 1, n, n); + for (size_t i = 0; i < n; ++i) { + ExpectSame(buffer[i], expected[i]); + } + } + } + } +} + +#if defined(MLAS_TARGET_RISCV64) +TEST(FusedActivation, ReluAndClipPreserveNaNPayloadsAndZero) { + constexpr uint32_t bits[] = {0x7fc12345, 0xffc12345, 0x7fa12345, 0xffa12345, + 0x80000000, 0, 1, 0x3f800000}; + for (auto kind : {MlasReluActivation, MlasClipActivation}) { + auto activation = MakeActivation(kind); + std::array expected; + std::array buffer; + for (size_t i = 0; i < expected.size(); ++i) expected[i] = bits[i % _countof(bits)]; + std::memcpy(buffer.data(), expected.data(), sizeof(buffer)); + MlasActivation(&activation, buffer.data(), nullptr, 1, buffer.size(), buffer.size()); + EXPECT_EQ(std::memcmp(buffer.data(), expected.data(), sizeof(buffer)), 0); + } +} +#endif + +TEST(FusedActivation, IdentityWithoutBiasIsNoOp) { + auto activation = MakeActivation(MlasIdentityActivation); + std::array bits = {0x7fc12345, 0xffc12345, 0x80000000, 0, 1, 0x80000001, + 0x7f800000, 0xff800000, 0x3f800000}; + std::array buffer; + std::memcpy(buffer.data(), bits.data(), sizeof(buffer)); + MlasActivation(&activation, buffer.data(), nullptr, 1, buffer.size(), buffer.size()); + EXPECT_EQ(std::memcmp(buffer.data(), bits.data(), sizeof(buffer)), 0); +} + +#if defined(MLAS_TARGET_RISCV64) && defined(MLAS_USE_RVV) +TEST(FusedActivationRvv, RuntimeDispatch) { + const auto& platform = GetMlasPlatform(); + if (platform.GemmFloatKernel == MlasGemmFloatKernelRvv) { + EXPECT_EQ(platform.ActivationRoutine, MlasActivationRvv); + } else { + EXPECT_EQ(platform.ActivationRoutine, nullptr); + } +} + +TEST(FusedActivationRvv, MatchesExistingRvvAndReferenceBoundaryAndRandomInputs) { + const auto& platform = GetMlasPlatform(); + if (platform.ActivationRoutine == nullptr) { + GTEST_SKIP() << "RVV is unavailable or forced off"; + } + ASSERT_EQ(platform.ActivationRoutine, MlasActivationRvv); + std::mt19937 generator(42); + std::uniform_real_distribution random(-10.0f, 10.0f); + const float exceptional[] = {-0.0f, 0.0f, -INFINITY, INFINITY, std::numeric_limits::quiet_NaN()}; + const float alphas[] = {0.2f, -0.5f, 0.0f, INFINITY, std::numeric_limits::quiet_NaN()}; + const float bias[] = {-0.0f, 0.5f, std::numeric_limits::quiet_NaN()}; + std::array actual; + std::array expected; + for (auto kind : kKinds) { + for (float alpha : alphas) { + auto activation = MakeActivation(kind); + if (kind == MlasLeakyReluActivation) activation.Parameters.LeakyRelu.alpha = alpha; + for (size_t n : {size_t(1), size_t(4), size_t(7), size_t(8), size_t(9), size_t(31), size_t(33), size_t(129)}) { + for (bool add_bias : {false, true}) { + SCOPED_TRACE(::testing::Message() << "kind=" << kind << " N=" << n << " alpha=" << alpha + << " bias=" << add_bias); + const size_t ldc = n + 3; + for (size_t i = 0; i < actual.size(); ++i) { + actual[i] = i % 2 == 0 ? exceptional[(i / 2) % _countof(exceptional)] : random(generator); + } + expected = actual; + if (kind == MlasHardSwishActivation) { + // The existing RVV routine leaves HardSwish to the generic path. + for (size_t row = 0; row < 3; ++row) { + for (size_t col = 0; col < n; ++col) { + const size_t index = row * ldc + col; + const float value = add_bias ? expected[index] + bias[row] : expected[index]; + expected[index] = Reference(activation, value); + } + } + } else { + ASSERT_TRUE(MlasActivationRvv(&activation, expected.data(), add_bias ? bias : nullptr, 3, n, ldc)); + } + MlasActivation(&activation, actual.data(), add_bias ? bias : nullptr, 3, n, ldc); + for (size_t i = 0; i < actual.size(); ++i) { + ExpectSame(actual[i], expected[i]); + } + } + } + } + } +} + +TEST(FusedActivationRvv, HardSigmoidNonfiniteParameters) { + const auto& platform = GetMlasPlatform(); + if (platform.ActivationRoutine == nullptr) { + GTEST_SKIP() << "RVV is unavailable or forced off"; + } + const float parameters[][2] = {{std::numeric_limits::max(), -INFINITY}, + {INFINITY, -INFINITY}, + {0.0f, NAN}, + {NAN, 0.0f}, + {1.0f, INFINITY}}; + for (const auto& parameter : parameters) { + auto activation = MakeActivation(MlasHardSigmoidActivation); + activation.Parameters.HardSigmoid.alpha = parameter[0]; + activation.Parameters.HardSigmoid.beta = parameter[1]; + std::array actual; + actual.fill(2.0f); + auto expected = actual; + EXPECT_FALSE(MlasFusedActivationRvv(&activation, actual.data(), nullptr, 1, actual.size(), actual.size())); + EXPECT_EQ(actual, expected); + ASSERT_TRUE(MlasActivationRvv(&activation, expected.data(), nullptr, 1, expected.size(), expected.size())); + MlasActivation(&activation, actual.data(), nullptr, 1, actual.size(), actual.size()); + for (size_t i = 0; i < actual.size(); ++i) ExpectSame(actual[i], expected[i]); + } +} + +TEST(FusedActivationRvv, UnsupportedActivationLeavesBufferUntouched) { + if (GetMlasPlatform().ActivationRoutine == nullptr) { + GTEST_SKIP() << "RVV is unavailable or forced off"; + } + for (auto kind : {MlasTanhActivation, MlasLogisticActivation}) { + auto activation = MakeActivation(kind); + std::array buffer = {-4.0f, -1.0f, 0.0f, 1.0f, 4.0f}; + const auto original = buffer; + const float bias = 1.0f; + EXPECT_FALSE(MlasFusedActivationRvv(&activation, buffer.data(), &bias, 1, buffer.size(), buffer.size())); + EXPECT_EQ(buffer, original); + } +} +#endif + +} // namespace