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
19 changes: 13 additions & 6 deletions onnxruntime/core/mlas/lib/activate.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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)) {
Comment on lines +496 to +498
return;
}
#endif
// Keep the existing RVV routine for short rows and unsupported kinds.
if (activation_routine(Activation, Buffer, Bias, M, N, ldc)) {
return;
Comment on lines +502 to +504
}
}
#endif

Expand Down
1 change: 1 addition & 0 deletions onnxruntime/core/mlas/lib/mlasi.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
160 changes: 158 additions & 2 deletions onnxruntime/core/mlas/lib/riscv64/activation_kernel_rvv.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -107,6 +107,162 @@ constexpr float ERF_A5 = 1.061405429f;

} // namespace

namespace {

template<MLAS_ACTIVATION_KIND ActivationKind, bool AddBias>
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<MLAS_ACTIVATION_KIND ActivationKind>
void
MlasFusedActivationDispatchRvv(
const MLAS_ACTIVATION* Activation,
float* Buffer,
const float* Bias,
size_t M,
size_t N,
size_t ldc
)
{
if (Bias != nullptr) {
MlasFusedActivationKernelRvv<ActivationKind, true>(Activation, Buffer, Bias, M, N, ldc);
} else if constexpr (ActivationKind != MlasIdentityActivation) {
MlasFusedActivationKernelRvv<ActivationKind, false>(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<MlasIdentityActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasReluActivation:
MlasFusedActivationDispatchRvv<MlasReluActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasLeakyReluActivation:
MlasFusedActivationDispatchRvv<MlasLeakyReluActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasClipActivation:
MlasFusedActivationDispatchRvv<MlasClipActivation>(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<MlasHardSigmoidActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasHardSwishActivation:
MlasFusedActivationDispatchRvv<MlasHardSwishActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
default:
return false;
}
}

extern "C"
void
MLASCALL
Expand Down
Loading
Loading