mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:20:33 +02:00
Gate mmvf_f8x2_e4m3_to_bf162 on FP8 availability, i.e. CTK >= 11.8
This commit is contained in:
@@ -8,9 +8,9 @@
|
|||||||
template <typename T>
|
template <typename T>
|
||||||
using mmvf_y_t = std::conditional_t<std::is_same_v<T, ggml_fp8_e4m3_t>, nv_bfloat16, float>;
|
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) {
|
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));
|
return static_cast<nv_bfloat162>(__nv_cvt_fp8x2_to_bf162raw(x, __NV_E4M3));
|
||||||
#else
|
#else
|
||||||
return make_bfloat162(
|
return make_bfloat162(
|
||||||
@@ -333,7 +333,7 @@ static __global__ void mul_mat_vec_f(
|
|||||||
#endif
|
#endif
|
||||||
} else if constexpr (std::is_same_v<T, ggml_fp8_e4m3_t>) {
|
} else if constexpr (std::is_same_v<T, ggml_fp8_e4m3_t>) {
|
||||||
const nv_bfloat162 * y2 = (const nv_bfloat162 *) y;
|
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 * 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);
|
const __nv_fp8x2_storage_t * gate_x2 = reinterpret_cast<const __nv_fp8x2_storage_t *>(gate_x);
|
||||||
nv_bfloat162 sum_bf[ncols_dst] = {};
|
nv_bfloat162 sum_bf[ncols_dst] = {};
|
||||||
|
|||||||
Reference in New Issue
Block a user