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
20 changes: 20 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,26 @@ else()
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -O3")
endif()
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} ${CMAKE_C_FLAGS}")

# RISC-V: enable the vector and half-precision extensions available on the
# target CPU. Prefer zfhmin/zvfhmin (native FP16 conversion) and fall back to
# the base vector extension when the toolchain does not accept the extras.
if(CMAKE_SYSTEM_PROCESSOR MATCHES "^riscv64" OR CMAKE_SYSTEM_PROCESSOR MATCHES "^rv64")
include(CheckCCompilerFlag)
check_cxx_compiler_flag("-march=rv64gcv_zfhmin_zvfhmin" MXNET_RISCV_HAS_GCV_ZFHMIN)
check_c_compiler_flag("-march=rv64gcv_zfhmin_zvfhmin" MXNET_RISCV_HAS_GCV_ZFHMIN_C)
if(MXNET_RISCV_HAS_GCV_ZFHMIN AND MXNET_RISCV_HAS_GCV_ZFHMIN_C)
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -march=rv64gcv_zfhmin_zvfhmin")
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -march=rv64gcv_zfhmin_zvfhmin")
else()
check_cxx_compiler_flag("-march=rv64gcv" MXNET_RISCV_HAS_GCV)
check_c_compiler_flag("-march=rv64gcv" MXNET_RISCV_HAS_GCV_C)
if(MXNET_RISCV_HAS_GCV AND MXNET_RISCV_HAS_GCV_C)
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -march=rv64gcv")
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -march=rv64gcv")
endif()
endif()
endif()
endif()

if(NOT mxnet_LINKER_LIBS)
Expand Down
158 changes: 158 additions & 0 deletions src/operator/nn/softmax-inl.h
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,10 @@
#include "../operator_common.h"
#include "../tensor/broadcast_reduce_op.h"

#if defined(__riscv) && defined(__riscv_v)
#include <riscv_vector.h>
#endif

using mshadow::red::limits::MinValue;

namespace mxnet {
Expand Down Expand Up @@ -64,6 +68,143 @@ struct log_softmax_fwd {
}
};

#if defined(__riscv) && defined(__riscv_v)
/*!
* \brief RVV 1.0 helpers for the CPU softmin (negate == true) forward path.
*
* The scalar reference performs a 3-pass cross-lane normalization
* (max -> exp(x - max) -> normalize) with per-element scalar libm expf calls,
* a scalar compare-branch max reduction and a per-element negation of the
* softmin input. These helpers keep the explicit three passes but operate on
* RVV vectors so the negation is applied once per vector (vfneg.v), the max and
* sum are produced by vector reductions, and exp is evaluated with a
* self-contained range-reduced polynomial.
*/
namespace softmax_rvv {

// Minimum row length for which the vectorized path is worthwhile.
constexpr index_t kMinRowLen = 8;

// Evaluate exp(x) for SEW=32 / LMUL=1 vectors through range reduction:
// n = rint(x * log2(e)), r = x - n*ln2_hi - n*ln2_lo, exp(x) = exp(r)*2^n
// Both ln2 terms are scaled by the integer n so the reduction keeps full
// precision. Only the softmax domain x <= 0 (or NaN) is exercised in practice;
// x is clamped to the range where exp stays representable and NaN is
// propagated explicitly.
inline vfloat32m1_t VecExp(vfloat32m1_t x, size_t vl) {
const float log2e = 1.4426950408889634f;
const float ln2_hi = 0.693145751953125f;
const float ln2_lo = 1.428606765330187e-06f;

vfloat32m1_t xc = __riscv_vfmax_vv_f32m1(
x, __riscv_vfmv_v_f_f32m1(-87.33654475f, vl), vl);
xc = __riscv_vfmin_vv_f32m1(
xc, __riscv_vfmv_v_f_f32m1(88.722839f, vl), vl);

vfloat32m1_t t = __riscv_vfmul_vf_f32m1(xc, log2e, vl);
vint32m1_t n = __riscv_vfcvt_x_f_v_i32m1(t, vl); // round-to-nearest-even
n = __riscv_vmax_vx_i32m1(n, -126, vl);
n = __riscv_vmin_vx_i32m1(n, 127, vl);
vfloat32m1_t nf = __riscv_vfcvt_f_x_v_f32m1(n, vl);

// r = x - n*ln2_hi - n*ln2_lo (both terms scaled by the integer n).
vfloat32m1_t r = __riscv_vfnmsac_vf_f32m1(xc, ln2_hi, nf, vl);
r = __riscv_vfnmsac_vf_f32m1(r, ln2_lo, nf, vl);

// exp(r) on |r| <= ln2/2 via a degree-7 Horner polynomial.
vfloat32m1_t p = __riscv_vfmv_v_f_f32m1(1.0f / 5040.0f, vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f / 720.0f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f / 120.0f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f / 24.0f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f / 6.0f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(0.5f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f, vl), vl);
p = __riscv_vfmadd_vv_f32m1(p, r, __riscv_vfmv_v_f_f32m1(1.0f, vl), vl);

// 2^n from the biased exponent bits.
vint32m1_t bits = __riscv_vadd_vx_i32m1(n, 127, vl);
bits = __riscv_vsll_vx_i32m1(bits, 23, vl);
vfloat32m1_t scale = __riscv_vreinterpret_v_i32m1_f32m1(bits);
vfloat32m1_t res = __riscv_vfmul_vv_f32m1(p, scale, vl);

// exp(+inf) = +inf.
vbool32_t ovf = __riscv_vmfgt_vf_f32m1_b32(x, 88.722839f, vl);
res = __riscv_vmerge_vvm_f32m1(
res, __riscv_vfmv_v_f_f32m1(__builtin_inff(), vl), ovf, vl);
// NaN in -> NaN out (the clamp above would otherwise hide it).
vbool32_t isnan = __riscv_vmfne_vv_f32m1_b32(x, x, vl);
res = __riscv_vmerge_vvm_f32m1(res, x, isnan, vl);
return res;
}

// Vectorized softmin (negate == true) forward over one contiguous float32 row.
template <typename OType>
inline void SoftminRow(const float* in, OType* out, size_t M) {
const size_t vlmax = __riscv_vsetvlmax_e32m1();
const float neg_inf = -__builtin_inff();

// Pass 1: mmax = max_j(-in[j]).
vfloat32m1_t vmax = __riscv_vfmv_v_f_f32m1(neg_inf, vlmax);
size_t max_j = 0;
for (; M - max_j >= vlmax; max_j += vlmax) {
vfloat32m1_t v = __riscv_vle32_v_f32m1(in + max_j, vlmax);
v = __riscv_vfneg_v_f32m1(v, vlmax); // softmin sign, once per vector
vmax = __riscv_vfmax_vv_f32m1(vmax, v, vlmax);
}
float mmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m1_f32m1(
vmax, __riscv_vfmv_v_f_f32m1(neg_inf, vlmax), vlmax));
if (max_j < M) {
const size_t vl = M - max_j;
vfloat32m1_t v = __riscv_vle32_v_f32m1(in + max_j, vl);
v = __riscv_vfneg_v_f32m1(v, vl);
mmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m1_f32m1(
v, __riscv_vfmv_v_f_f32m1(mmax, vl), vl));
}

// Pass 2: sum = sum_j exp(-in[j] - mmax).
vfloat32m1_t vsum = __riscv_vfmv_v_f_f32m1(0.0f, vlmax);
size_t sum_j = 0;
for (; M - sum_j >= vlmax; sum_j += vlmax) {
vfloat32m1_t v = __riscv_vle32_v_f32m1(in + sum_j, vlmax);
v = __riscv_vfneg_v_f32m1(v, vlmax);
v = __riscv_vfsub_vf_f32m1(v, mmax, vlmax);
vsum = __riscv_vfadd_vv_f32m1(vsum, VecExp(v, vlmax), vlmax);
}
float sum = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredusum_vs_f32m1_f32m1(
vsum, __riscv_vfmv_v_f_f32m1(0.0f, vlmax), vlmax));
if (sum_j < M) {
const size_t vl = M - sum_j;
vfloat32m1_t v = __riscv_vle32_v_f32m1(in + sum_j, vl);
v = __riscv_vfneg_v_f32m1(v, vl);
v = __riscv_vfsub_vf_f32m1(v, mmax, vl);
vfloat32m1_t e = VecExp(v, vl);
sum = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredusum_vs_f32m1_f32m1(
e, __riscv_vfmv_v_f_f32m1(sum, vl), vl));
}

// Pass 3: out[j] = exp(-in[j] - mmax) / sum.
for (size_t j = 0; j < M;) {
const size_t vl = __riscv_vsetvl_e32m1(M - j);
vfloat32m1_t v = __riscv_vle32_v_f32m1(in + j, vl);
v = __riscv_vfneg_v_f32m1(v, vl);
v = __riscv_vfsub_vf_f32m1(v, mmax, vl);
vfloat32m1_t res = __riscv_vfdiv_vf_f32m1(VecExp(v, vl), sum, vl);
if constexpr (std::is_same<OType, float>::value) {
__riscv_vse32_v_f32m1(reinterpret_cast<float*>(out + j), res, vl);
} else {
#if defined(__riscv_zvfhmin)
vfloat16mf2_t h = __riscv_vfncvt_f_f_w_f16mf2(res, vl);
vuint16mf2_t hb = __riscv_vreinterpret_v_f16mf2_u16mf2(h);
__riscv_vse16_v_u16mf2(reinterpret_cast<uint16_t*>(out + j), hb, vl);
#endif
}
j += vl;
}
}

} // namespace softmax_rvv
#endif // defined(__riscv) && defined(__riscv_v)

template <typename OP,
bool negate,
typename AType,
Expand Down Expand Up @@ -92,6 +233,23 @@ inline void Softmax(Stream<cpu>* s,
for (index_t i = 0; i < N; ++i) {
index_t base = unravel_dot(i, sshape, stride);

#if defined(__riscv) && defined(__riscv_v)
// Vectorized softmin forward for contiguous float32 rows.
if constexpr (std::is_same<OP, softmax_fwd>::value &&
std::is_same<DType, float>::value &&
(std::is_same<OType, float>::value
#if defined(__riscv_zvfhmin)
|| std::is_same<OType, mshadow::half::half_t>::value
#endif
)) {
if (negate && temperature == static_cast<DType>(1.0) && sa == 1 &&
M >= softmax_rvv::kMinRowLen) {
softmax_rvv::SoftminRow(in + base, out + base, static_cast<size_t>(M));
continue;
}
}
#endif

DType mmax = negate ? -in[base] : in[base];
DType val;
for (index_t j = 1; j < M; ++j) {
Expand Down