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
11 changes: 11 additions & 0 deletions onnxruntime/core/mlas/lib/activate.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -489,6 +489,17 @@ Return Value:

--*/
{
#if defined(MLAS_TARGET_RISCV64) && defined(MLAS_USE_RVV) && \
!defined(FORCE_GENERIC_ALGORITHMS)
// Short rows do not amortize the dispatch and vector setup costs.
if (N >= 32 && (Activation->ActivationKind != MlasIdentityActivation || Bias != nullptr)) {
const auto activation_override = GetMlasPlatform().MlasActivationOverride;
if (activation_override != nullptr && activation_override(Activation, Buffer, Bias, M, N, ldc)) {
return;
}
}
#endif

switch (Activation->ActivationKind) {

case MlasIdentityActivation:
Expand Down
14 changes: 14 additions & 0 deletions onnxruntime/core/mlas/lib/mlasi.h
Original file line number Diff line number Diff line change
Expand Up @@ -651,6 +651,18 @@ void
size_t OutputCountRightPad
);

// Return false to request the generic fallback without modifying Buffer.
typedef
bool
(MLASCALL MLAS_ACTIVATION_OVERRIDE)(
const MLAS_ACTIVATION* Activation,
float* Buffer,
const float* Bias,
size_t M,
size_t N,
size_t ldc
);

typedef
void
(MLASCALL MLAS_COMPUTE_UNARY_FLOAT_KERNEL)(
Expand Down Expand Up @@ -1450,6 +1462,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_OVERRIDE MlasActivationRvv;
MLAS_COMPUTE_SUMEXP_FLOAT_KERNEL MlasComputeSumExpF32KernelRvv;
MLAS_REDUCE_MAXIMUM_FLOAT_KERNEL MlasReduceMaximumF32KernelRvv;
MLAS_COMPUTE_SOFTMAX_OUTPUT_FLOAT_KERNEL MlasComputeSoftmaxOutputF32KernelRvv;
Expand Down Expand Up @@ -1887,6 +1900,7 @@ struct MLAS_PLATFORM {
#endif

#if defined(MLAS_TARGET_RISCV64) && defined(MLAS_USE_RVV)
MLAS_ACTIVATION_OVERRIDE* MlasActivationOverride{nullptr};
MLAS_CONV_FLOAT_KERNEL* ConvNchwFloatKernel;
MLAS_CONV_FLOAT_KERNEL* ConvNchwcFloatKernel;
MLAS_CONV_DEPTHWISE_FLOAT_KERNEL* ConvDepthwiseFloatKernel;
Expand Down
1 change: 1 addition & 0 deletions onnxruntime/core/mlas/lib/platform.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -354,6 +354,7 @@ Return Value:
has_rvv = false;
}
if (has_rvv) {
this->MlasActivationOverride = MlasActivationRvv;
this->GemmFloatKernel = MlasGemmFloatKernelRvv;
this->GemmU8S8Dispatch = &MlasGemmQuantDispatchRvv;
this->GemmU8U8Dispatch = &MlasGemmQuantDispatchRvv;
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
MlasFusedActivationRvv(
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
MlasActivationRvv(
const MLAS_ACTIVATION* Activation,
float* Buffer,
const float* Bias,
size_t M,
size_t N,
size_t ldc
)
{
switch (Activation->ActivationKind) {
case MlasIdentityActivation:
MlasFusedActivationRvv<MlasIdentityActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasReluActivation:
MlasFusedActivationRvv<MlasReluActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasLeakyReluActivation:
MlasFusedActivationRvv<MlasLeakyReluActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasClipActivation:
MlasFusedActivationRvv<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;
}
MlasFusedActivationRvv<MlasHardSigmoidActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
case MlasHardSwishActivation:
MlasFusedActivationRvv<MlasHardSwishActivation>(Activation, Buffer, Bias, M, N, ldc);
return true;
default:
return false;
}
}

extern "C"
void
MLASCALL
Expand Down
Loading
Loading