mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 02:47:26 -05:00
Gate mmvf_f8x2_e4m3_to_bf162 on FP8 availability, i.e. CTK >= 11.8
This commit is contained in:
@@ -7,9 +7,9 @@
|
||||
template <typename T>
|
||||
using mmvf_y_t = std::conditional_t<std::is_same_v<T, ggml_fp8_e4m3_t>, nv_bfloat16, float>;
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
#if defined(FP8_AVAILABLE) && !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
static __device__ __forceinline__ nv_bfloat162 mmvf_f8x2_e4m3_to_bf162(__nv_fp8x2_storage_t x) {
|
||||
#if defined(FP8_AVAILABLE) && CUDART_VERSION >= 13020
|
||||
#if CUDART_VERSION >= 13020
|
||||
return static_cast<nv_bfloat162>(__nv_cvt_fp8x2_to_bf162raw(x, __NV_E4M3));
|
||||
#else
|
||||
return make_bfloat162(
|
||||
@@ -332,7 +332,7 @@ static __global__ void mul_mat_vec_f(
|
||||
#endif
|
||||
} else if constexpr (std::is_same_v<T, ggml_fp8_e4m3_t>) {
|
||||
const nv_bfloat162 * y2 = (const nv_bfloat162 *) y;
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
#if defined(FP8_AVAILABLE) && !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
const __nv_fp8x2_storage_t * x2 = reinterpret_cast<const __nv_fp8x2_storage_t *>(x);
|
||||
const __nv_fp8x2_storage_t * gate_x2 = reinterpret_cast<const __nv_fp8x2_storage_t *>(gate_x);
|
||||
nv_bfloat162 sum_bf[ncols_dst] = {};
|
||||
|
||||
Reference in New Issue
Block a user