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 @@ -241,6 +241,7 @@ if(MSVC)
set(CMAKE_C_FLAGS "/MP")
set(CMAKE_CXX_FLAGS "${CMAKE_C_FLAGS} ${CMAKE_CXX_FLAGS} /bigobj")
else()
include(CheckCCompilerFlag)
include(CheckCXXCompilerFlag)
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -Wall -Wno-sign-compare")
if(CMAKE_BUILD_TYPE STREQUAL "Debug")
Expand All @@ -254,6 +255,25 @@ else()
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -O3")
endif()
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} ${CMAKE_C_FLAGS}")

# RISC-V: expose the vector extension plus half-precision conversions so the
# RVV kernels and native fcvt.h.s / vfncvt.f.f.v paths are available. Fall
# back to plain rv64gcv when zfhmin/zvfhmin are rejected by the toolchain.
if(CMAKE_SYSTEM_PROCESSOR MATCHES "riscv64")
check_c_compiler_flag("-march=rv64gcv_zfhmin_zvfhmin" MXNET_RISCV_ZFHMIN_C_SUPPORTED)
check_cxx_compiler_flag("-march=rv64gcv_zfhmin_zvfhmin" MXNET_RISCV_ZFHMIN_CXX_SUPPORTED)
if(MXNET_RISCV_ZFHMIN_C_SUPPORTED AND MXNET_RISCV_ZFHMIN_CXX_SUPPORTED)
set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -march=rv64gcv_zfhmin_zvfhmin")
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -march=rv64gcv_zfhmin_zvfhmin")
else()
check_c_compiler_flag("-march=rv64gcv" MXNET_RISCV_GV_C_SUPPORTED)
check_cxx_compiler_flag("-march=rv64gcv" MXNET_RISCV_GV_CXX_SUPPORTED)
if(MXNET_RISCV_GV_C_SUPPORTED AND MXNET_RISCV_GV_CXX_SUPPORTED)
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
175 changes: 175 additions & 0 deletions src/operator/nn/softmax-inl.h
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,16 @@
#include "../operator_common.h"
#include "../tensor/broadcast_reduce_op.h"

#if defined(__riscv) && defined(__riscv_v)
#include <riscv_vector.h>
// Zvfh/Zvfhmin expose the vector FP16 narrowing conversion (vfncvt.f.f.v).
#if defined(__riscv_zvfh) || defined(__riscv_zvfhmin)
#define MXNET_RVV_F16_NARROW 1
#else
#define MXNET_RVV_F16_NARROW 0
#endif
#endif

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

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

#if defined(__riscv) && defined(__riscv_v)
namespace softmax_rvv_detail {

/*!
* \brief Vectorized exp(x) for an f32m1 vector.
*
* Range-reduced: k = round(x * log2(e)), r = x - k*ln2,
* exp(x) = 2^k * (1 + r*Q(r)) with Q a degree-5 polynomial. 2^k is rebuilt from
* the IEEE-754 exponent field. Finite inputs are clamped so the exponent field
* cannot overflow/underflow; NaN lanes are re-injected to preserve propagation.
*/
inline vfloat32m1_t vexp_f32m1(vfloat32m1_t x, size_t vl) {
const vbool32_t is_nan = __riscv_vmfne_vv_f32m1_b32(x, x, vl);

// ln(FLT_MIN) and 128*ln(2): keeps the reconstructed exponent in normal range.
vfloat32m1_t xc = __riscv_vfmax_vf_f32m1(x, -87.3365478515625f, vl);
xc = __riscv_vfmin_vf_f32m1(xc, 88.72283935546875f, vl);

const float log2e = 1.4426950408889634f;
const vfloat32m1_t kf = __riscv_vfmul_vf_f32m1(xc, log2e, vl);
const vint32m1_t ki = __riscv_vfcvt_x_f_v_i32m1(kf, vl); // round-to-nearest
const vfloat32m1_t k = __riscv_vfcvt_f_x_v_f32m1(ki, vl);

// r = xc - k*ln2, with ln2 split into hi+lo and BOTH parts scaled by k to
// reduce cancellation error: r = xc - k*ln2_hi - k*ln2_lo.
vfloat32m1_t r =
__riscv_vfsub_vv_f32m1(xc, __riscv_vfmul_vf_f32m1(k, 0.6931457519531250f, vl), vl);
r = __riscv_vfsub_vv_f32m1(r, __riscv_vfmul_vf_f32m1(k, 1.428606765330187e-06f, vl), vl);

// Q(r) = 1 + r/2 + r^2/6 + r^3/24 + r^4/120 + r^5/720 (Horner).
vfloat32m1_t q = __riscv_vfmv_v_f_f32m1(1.0f / 720.0f, vl);
q = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(q, r, vl), 1.0f / 120.0f, vl);
q = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(q, r, vl), 1.0f / 24.0f, vl);
q = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(q, r, vl), 1.0f / 6.0f, vl);
q = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(q, r, vl), 0.5f, vl);
q = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(q, r, vl), 1.0f, vl);
const vfloat32m1_t er = __riscv_vfadd_vf_f32m1(__riscv_vfmul_vv_f32m1(r, q, vl), 1.0f, vl);

// 2^k == ((k + 127) << 23) interpreted as a float bit pattern.
const vint32m1_t bits = __riscv_vsll_vx_i32m1(__riscv_vadd_vx_i32m1(ki, 127, vl), 23, vl);
const vfloat32m1_t scale = __riscv_vreinterpret_v_i32m1_f32m1(bits);

const vfloat32m1_t res = __riscv_vfmul_vv_f32m1(er, scale, vl);
return __riscv_vmerge_vvm_f32m1(res, x, is_nan, vl); // NaN in -> NaN out
}

/*!
* \brief Store an f32m1 vector into OType storage (float direct, half_t narrowed).
*/
template <typename OType>
inline void vstore_as(OType* out, vfloat32m1_t v, size_t vl) {
if constexpr (std::is_same<OType, float>::value) {
__riscv_vse32_v_f32m1(out, v, vl);
} else {
#if MXNET_RVV_F16_NARROW
// vfncvt.f.f.v uses the dynamic rounding mode frm; its reset default is RNE,
// matching mshadow float2half's hardcoded RNE (MSHADOW_HALF_ROUND_TO_NEAREST).
// Callers that change frm at runtime or set MSHADOW_HALF_ROUND_TO_NEAREST=0
// would diverge and must not enable this fast path. half_t is a 16-bit
// wrapper over uint16_t, so the destination reinterpret is layout-preserving.
const vfloat16mf2_t h = __riscv_vfncvt_f_f_w_f16mf2(v, vl);
const vuint16mf2_t bits = __riscv_vreinterpret_v_f16mf2_u16mf2(h);
__riscv_vse16_v_u16mf2(reinterpret_cast<uint16_t*>(out), bits, vl);
#else
// Portable fallback: narrow element by element (only reached when the
// toolchain lacks Zvfhmin, in which case half_t output stays on the scalar
// path and this is not instantiated).
for (size_t i = 0; i < vl; ++i) {
const float lane = __riscv_vfmv_f_s_f32m1_f32(__riscv_vslidedown_vx_f32m1(v, i, vl));
out[i] = OType(lane);
}
#endif
}
}

/*!
* \brief RVV log-softmax forward for one contiguous float32 row.
*
* Mirrors the reference scalar semantics: stable max subtraction, exp-sum with
* the same shift, then out = (x - max) - log(sum). Block partial sums are
* accumulated in double to match the AType=double scalar contract; only the
* intra-block f32 horizontal reduction is unordered, spanning at most
* VLEN/32 lanes.
*/
template <typename OType>
inline void log_softmax_fwd_row(const float* in, OType* out, index_t M) {
// Pass 1: max reduction (scalar seed carried across blocks).
float mmax;
{
size_t vl = __riscv_vsetvl_e32m1(static_cast<size_t>(M));
vfloat32m1_t seed = __riscv_vfmv_v_f_f32m1(-INFINITY, vl);
for (index_t j = 0; j < M; j += static_cast<index_t>(vl)) {
const size_t cur = __riscv_vsetvl_e32m1(static_cast<size_t>(M - j));
const vfloat32m1_t x = __riscv_vle32_v_f32m1(in + j, cur);
seed = __riscv_vfredmax_vs_f32m1_f32m1(x, seed, cur);
vl = cur;
}
mmax = __riscv_vfmv_f_s_f32m1_f32(seed);
}

// Pass 2: sum of exp(x - max). Each block is reduced horizontally and the
// block partials are accumulated in double, matching the AType=double scalar
// reference (which accumulates sum in AType) and keeping log(sum) accurate
// for long axes. The intra-block f32 reduction is unordered (vfredusum) but
// spans at most VLEN/32 lanes, a negligible deviation from the reference.
double sum = 0.0;
{
size_t vl = __riscv_vsetvl_e32m1(static_cast<size_t>(M));
const vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, vl);
for (index_t j = 0; j < M; j += static_cast<index_t>(vl)) {
const size_t cur = __riscv_vsetvl_e32m1(static_cast<size_t>(M - j));
const vfloat32m1_t x = __riscv_vle32_v_f32m1(in + j, cur);
const vfloat32m1_t d = __riscv_vfsub_vf_f32m1(x, mmax, cur);
const vfloat32m1_t e = vexp_f32m1(d, cur);
const vfloat32m1_t blk = __riscv_vfredusum_vs_f32m1_f32m1(e, zero, cur);
sum += static_cast<double>(__riscv_vfmv_f_s_f32m1_f32(blk));
vl = cur;
}
}

// Pass 3: out = (x - max) - log(sum); log(sum) is evaluated once in double,
// as in the AType=double scalar reference (log_softmax_fwd::Map(double)).
const float lsum = static_cast<float>(std::log(sum));
{
size_t vl = __riscv_vsetvl_e32m1(static_cast<size_t>(M));
for (index_t j = 0; j < M; j += static_cast<index_t>(vl)) {
const size_t cur = __riscv_vsetvl_e32m1(static_cast<size_t>(M - j));
const vfloat32m1_t x = __riscv_vle32_v_f32m1(in + j, cur);
vfloat32m1_t d = __riscv_vfsub_vf_f32m1(x, mmax, cur);
d = __riscv_vfsub_vf_f32m1(d, lsum, cur);
vstore_as<OType>(out + j, d, cur);
vl = cur;
}
}
}

/*!
* \brief Compile-time eligibility for the RVV log-softmax fast path.
*/
template <typename OP, bool negate, typename DType, typename OType>
struct rvv_log_softmax_fwd_eligible {
static constexpr bool value =
std::is_same<OP, log_softmax_fwd>::value && !negate &&
std::is_same<DType, float>::value &&
(std::is_same<OType, float>::value ||
(std::is_same<OType, mshadow::half::half_t>::value && MXNET_RVV_F16_NARROW == 1));
};

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

template <typename OP,
bool negate,
typename AType,
Expand All @@ -87,6 +248,20 @@ inline void Softmax(Stream<cpu>* s,
sshape[axis] = 1;
index_t sa = stride[axis];

#if defined(__riscv) && defined(__riscv_v)
// RVV fast path for contiguous float32 log-softmax forward at temperature 1.
if constexpr (softmax_rvv_detail::rvv_log_softmax_fwd_eligible<OP, negate, DType, OType>::value) {
if (length == nullptr && sa == 1 && temperature == DType(1.0)) {
#pragma omp parallel for
for (index_t i = 0; i < N; ++i) {
index_t base = unravel_dot(i, sshape, stride);
softmax_rvv_detail::log_softmax_fwd_row<OType>(in + base, out + base, M);
}
return;
}
}
#endif

if (length == nullptr) {
#pragma omp parallel for
for (index_t i = 0; i < N; ++i) {
Expand Down