diff --git a/CMakeLists.txt b/CMakeLists.txt index 53a6978f45..cd12328b3d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -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) diff --git a/src/operator/nn/softmax-inl.h b/src/operator/nn/softmax-inl.h index 3023cd7177..20de57d2ee 100644 --- a/src/operator/nn/softmax-inl.h +++ b/src/operator/nn/softmax-inl.h @@ -34,6 +34,10 @@ #include "../operator_common.h" #include "../tensor/broadcast_reduce_op.h" +#if defined(__riscv) && defined(__riscv_v) +#include +#endif + using mshadow::red::limits::MinValue; namespace mxnet { @@ -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 +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::value) { + __riscv_vse32_v_f32m1(reinterpret_cast(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(out + j), hb, vl); +#endif + } + j += vl; + } +} + +} // namespace softmax_rvv +#endif // defined(__riscv) && defined(__riscv_v) + template * 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::value && + std::is_same::value && + (std::is_same::value +#if defined(__riscv_zvfhmin) + || std::is_same::value +#endif + )) { + if (negate && temperature == static_cast(1.0) && sa == 1 && + M >= softmax_rvv::kMinRowLen) { + softmax_rvv::SoftminRow(in + base, out + base, static_cast(M)); + continue; + } + } +#endif + DType mmax = negate ? -in[base] : in[base]; DType val; for (index_t j = 1; j < M; ++j) {