namespace syclexp = sycl::ext::oneapi::experimental;
+#if defined(__INTEL_LLVM_COMPILER) && __has_include(<sycl/ext/oneapi/bfloat16.hpp>)
+ #include <sycl/ext/oneapi/bfloat16.hpp>
+ #ifndef GGML_SYCL_HAS_BF16
+ #define GGML_SYCL_HAS_BF16
+ #endif
+#endif
+
#if GGML_SYCL_DNNL
#include "dnnl.hpp"
#include "dnnl_sycl.hpp"
#include "dequantize.hpp"
#include "presets.hpp"
-#if defined(__INTEL_LLVM_COMPILER)
- #if __has_include(<sycl/ext/oneapi/bfloat16.hpp>)
- #include <sycl/ext/oneapi/bfloat16.hpp>
- #define GGML_SYCL_HAS_BF16
- #endif
-#endif
-
template <int qk, int qr, dequantize_kernel_t dequantize_kernel, typename dst_t>
static void dequantize_block(const void * __restrict__ vx, dst_t * __restrict__ y, const int64_t k,
const sycl::nd_item<3> &item_ct1) {
}
+#ifdef GGML_SYCL_HAS_BF16
+to_bf16_sycl_t ggml_get_to_bf16_sycl(ggml_type type, ggml_tensor * /*dst*/) {
+ switch (type) {
+ case GGML_TYPE_F32:
+ return convert_unary_sycl<float>;
+ case GGML_TYPE_F16:
+ return convert_unary_sycl<sycl::half>;
+ case GGML_TYPE_BF16:
+ return convert_unary_sycl<sycl::ext::oneapi::bfloat16>;
+ default:
+ GGML_ABORT("fatal error: unsupport data type=%s\n", ggml_type_name(type));
+ return nullptr;
+ }
+}
+#endif
+
to_fp16_nc_sycl_t ggml_get_to_fp16_nc_sycl(ggml_type type) {
switch (type) {
case GGML_TYPE_F32:
to_fp16_sycl_t ggml_get_to_fp16_sycl(ggml_type type, ggml_tensor * dst);
to_fp32_sycl_t ggml_get_to_fp32_sycl(ggml_type type, ggml_tensor * dst);
+#ifdef GGML_SYCL_HAS_BF16
+typedef to_t_sycl_t<sycl::ext::oneapi::bfloat16> to_bf16_sycl_t;
+to_bf16_sycl_t ggml_get_to_bf16_sycl(ggml_type type, ggml_tensor * dst);
+#endif
+
// Nc = Non-contiguous
template <typename T>
using to_t_nc_sycl_t = void (*)(const void * x, T * y, int64_t ne00, int64_t ne01, int64_t ne02, int64_t ne03,
inline dst_t ggml_sycl_cast(src_t x) {
if constexpr (std::is_same_v<dst_t, src_t>) {
return x;
+#ifdef GGML_SYCL_HAS_BF16
} else if constexpr (std::is_same_v<dst_t, sycl::ext::oneapi::bfloat16>) {
return sycl::ext::oneapi::bfloat16(float(x));
} else if constexpr (std::is_same_v<src_t, sycl::ext::oneapi::bfloat16>) {
return static_cast<float>(x);
+#endif
} else if constexpr (std::is_same_v<src_t, sycl::float2> && std::is_same_v<dst_t, sycl::half2>) {
return x.template convert<sycl::half, sycl::rounding_mode::rte>();
+#ifdef GGML_SYCL_HAS_BF16
} else if constexpr (std::is_same_v<src_t, sycl::float2> &&
std::is_same_v<dst_t, sycl::vec<sycl::ext::oneapi::bfloat16, 2>>) {
return {x.x, x.y};
+#endif
} else if constexpr(std::is_same_v<dst_t, int32_t>) {
return int32_t(x);
} else {
static constexpr dt to_dt() {
if constexpr (std::is_same_v<T, float>) return dt::f32;
else if constexpr (std::is_same_v<T, sycl::half>) return dt::f16;
+#ifdef GGML_SYCL_HAS_BF16
+ else if constexpr (std::is_same_v<T, sycl::ext::oneapi::bfloat16>) return dt::bf16;
+#endif
else static_assert(0);
}
#else
bool use_fp16 = false;
#endif
+
+#if GGML_SYCL_DNNL && defined(GGML_SYCL_HAS_BF16)
+ // Fast path for bf16 src0
+ if (src0->type == GGML_TYPE_BF16 && !g_ggml_sycl_disable_dnn && ggml_is_contiguous(src0) &&
+ row_diff == src0->ne[1]) {
+ using bf16_t = sycl::ext::oneapi::bfloat16;
+ ggml_sycl_pool_alloc<bf16_t> src1_as_bf16(ctx.pool(), src1_ncols*ne10);
+ if (src1->type != GGML_TYPE_BF16) {
+ const to_bf16_sycl_t to_bf16_sycl = ggml_get_to_bf16_sycl(src1->type, dst);
+ GGML_ASSERT(to_bf16_sycl != nullptr);
+ to_bf16_sycl(src1_ddf_i, src1_as_bf16.get(), src1_ncols*ne10, stream);
+ } else {
+ stream->memcpy(src1_as_bf16.get(), src1_ddf_i, src1_ncols*ne10*sizeof(bf16_t));
+ }
+ DnnlGemmWrapper::row_gemm(ctx, row_diff, src1_ncols, ne10,
+ src0_dd_i, DnnlGemmWrapper::to_dt<bf16_t>(),
+ src1_as_bf16.get(), DnnlGemmWrapper::to_dt<bf16_t>(),
+ dst_dd_i, DnnlGemmWrapper::to_dt<float>(), stream);
+ GGML_UNUSED(dst);
+ GGML_UNUSED(src1_ddq_i);
+ GGML_UNUSED(src1_padded_row_size);
+ return;
+ }
+#endif
+
if ((src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type)) && use_fp16 && ggml_is_contiguous(src0) &&
row_diff == src0->ne[1] && dst->op_params[0] == GGML_PREC_DEFAULT) {
ggml_sycl_pool_alloc<sycl::half> src0_as_f16(ctx.pool());
}
}
} else {
- ggml_sycl_pool_alloc<char> src1_contiguous(ctx.pool(), sizeof(float)*ggml_nelements(src1));
- ggml_sycl_pool_alloc<char> dst_contiguous(ctx.pool(), sizeof(float)*ggml_nelements(dst));
+ const int64_t n_routed_rows = ids->ne[1] * n_ids;
+ ggml_sycl_pool_alloc<char> src1_contiguous(ctx.pool(), sizeof(float)*n_routed_rows*ne10);
+ ggml_sycl_pool_alloc<char> dst_contiguous(ctx.pool(), sizeof(float)*n_routed_rows*ne0);
src1_row.data = src1_contiguous.get();
dst_row.data = dst_contiguous.get();
namespace utils {
template<typename T>
static constexpr bool is_arithmetic_v() {
- return std::is_arithmetic_v<T> || std::is_same_v<T, sycl::half> || std::is_same_v<T, sycl::ext::oneapi::bfloat16>;
+ return std::is_arithmetic_v<T> || std::is_same_v<T, sycl::half>
+#ifdef GGML_SYCL_HAS_BF16
+ || std::is_same_v<T, sycl::ext::oneapi::bfloat16>
+#endif
+ ;
}
}
stream
);
break;
+#ifdef GGML_SYCL_HAS_BF16
case GGML_TYPE_BF16:
set_rows_sycl<TIn, TIdx, sycl::ext::oneapi::bfloat16>(
src0_d, src1_d, (char *)dst->data,
stream
);
break;
+#endif
case GGML_TYPE_Q8_0:
set_rows_sycl_q<TIdx, block_q8_0, QK8_0, cpy_blck_f32_q8_0>(src0_d, src1_d, (block_q8_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;