diff --git a/ggml/src/ggml-cuda/convert.cu b/ggml/src/ggml-cuda/convert.cu index 6de6c82b9b..1e72b4a498 100644 --- a/ggml/src/ggml-cuda/convert.cu +++ b/ggml/src/ggml-cuda/convert.cu @@ -29,6 +29,14 @@ static __global__ void dequantize_block(const void * __restrict__ vx, dst_t * __ const int64_t iybs = i00 - i00%qk; // y block start index const int64_t y_offset = qr == 1 ? 1 : qk/2; + if constexpr (qk == 1) { + if (i00 + 1 == ne00) { + const int64_t iy0 = (i0203*ne01 + i01)*ne00 + i00; + y[iy0] = ggml_cuda_cast(ggml_cuda_f8_e4m3_to_fp32(((const uint8_t *) vx)[ib])); + continue; + } + } + // dequantize float2 v; dequantize_kernel(vx, ib, iqs, v); diff --git a/ggml/src/ggml-cuda/getrows.cu b/ggml/src/ggml-cuda/getrows.cu index 71c3321d25..3d0a63b243 100644 --- a/ggml/src/ggml-cuda/getrows.cu +++ b/ggml/src/ggml-cuda/getrows.cu @@ -30,6 +30,13 @@ static __global__ void k_get_rows( const int iybs = i00 - i00%qk; // dst block start index const int y_offset = qr == 1 ? 1 : qk/2; + if constexpr (qk == 1) { + if (i00 + 1 == ne00) { + dst_row[i00] = ggml_cuda_cast(ggml_cuda_f8_e4m3_to_fp32(((const uint8_t *) src0_row)[i00])); + continue; + } + } + // dequantize float2 v; dequantize_kernel(src0_row, ib, iqs, v); @@ -178,7 +185,7 @@ static void get_rows_cuda_q( const size_t s12 = nb12 / sizeof(int32_t); // const size_t s13 = nb13 / sizeof(int32_t); - GGML_ASSERT(ne00 % 2 == 0); + GGML_ASSERT(qk == 1 || ne00 % 2 == 0); GGML_ASSERT(ne12 > 0); GGML_ASSERT(ne11 <= std::numeric_limits::max() / ne12); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index a4d2af86d1..bd9a034305 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9376,6 +9376,9 @@ static std::vector> make_test_cases_eval() { } test_cases.emplace_back(new test_get_rows(GGML_TYPE_F32, 1, 8, 2, 1, 1, false)); + for (int n : {1, 3, 257}) { + test_cases.emplace_back(new test_get_rows(GGML_TYPE_F8_E4M3, n, 5, 3, 1, 1, false)); + } for (ggml_type type : all_types) { for (int b : {1, 7}) { for (bool v : {false, true}) {