diff --git a/ggml/src/ggml-cuda/mmq-config-ampere.cuh b/ggml/src/ggml-cuda/mmq-config-ampere.cuh index 9261c67f65a3..ca46838be0eb 100644 --- a/ggml/src/ggml-cuda/mmq-config-ampere.cuh +++ b/ggml/src/ggml-cuda/mmq-config-ampere.cuh @@ -54,6 +54,7 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); + CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); @@ -61,6 +62,10 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); + CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); + CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); + CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); + CASE(GGML_TYPE_PTQ1_0, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, false); CASE(GGML_TYPE_Q4_0, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); CASE(GGML_TYPE_Q4_0, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, true, true); diff --git a/ggml/src/ggml-cuda/mmq-load-tiles.cuh b/ggml/src/ggml-cuda/mmq-load-tiles.cuh index 7f7192eac199..04d4493edf08 100644 --- a/ggml/src/ggml-cuda/mmq-load-tiles.cuh +++ b/ggml/src/ggml-cuda/mmq-load-tiles.cuh @@ -257,21 +257,6 @@ template static __device__ __forceinline_ } #if !defined(GGML_USE_HIP) -static __device__ -__forceinline__ void ggml_cuda_mmq_decode_ptq1_0_qs4(uint32_t packed, int * __restrict__ dst, int stride) { - uint32_t v_lo = __byte_perm(packed, 0, 0x4140); - uint32_t v_hi = __byte_perm(packed, 0, 0x4342); - -# pragma unroll - for (int t = 0; t < 5; ++t) { - const uint32_t w_lo = v_lo * 3; - const uint32_t w_hi = v_hi * 3; - v_lo = w_lo & 0x00FF00FF; - v_hi = w_hi & 0x00FF00FF; - dst[t * stride] = __vsub4(__byte_perm(w_lo, w_hi, 0x7531), 0x01010101); - } -} - template static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_ptq1_0(const char * __restrict__ x, int * __restrict__ x_tile, @@ -315,22 +300,31 @@ static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_ptq1_0(const cha int * row = x_qs + i * (2 * MMQ_TILE_NE_K + 1) + kbx * (QK_PTQ1_0 / 4); # endif - if (lane < 4) { - ggml_cuda_mmq_decode_ptq1_0_qs4(get_int_b4(bxi->qs, lane), row + lane, 4); - } else if (lane < 6) { - const int g = lane - 4; - ggml_cuda_mmq_decode_ptq1_0_qs4(get_int_b4(bxi->qs + 16, g), row + 20 + g, 2); - } else if (lane == 6) { - uint32_t v = (uint32_t) bxi->qh[0] | ((uint32_t) bxi->qh[1] << 16); + // Branch-free unpack: all 8 lanes run the same 5-iteration trit loop on their own 32-bit word. Only smem store offsets differ, so the warp does not diverge. + const uint32_t packed = get_int_b4(bxi->qs, lane < 7 ? lane : 6); + uint32_t v_lo = __byte_perm(packed, 0, 0x4140); // bytes 0,1 as 16-bit lanes + uint32_t v_hi = __byte_perm(packed, 0, 0x4342); // bytes 2,3 + const bool full_lane = lane < 6; + v_hi = full_lane ? v_hi : v_lo; // lane 6: both halves walk qh0/qh1 + const int dst_base = lane < 4 ? lane : 16 + lane; // lanes 4,5 -> 20,21 + const int dst_stride = lane < 4 ? 4 : 2; + int q[5]; # pragma unroll - for (int t = 0; t < 4; t += 2) { - const uint32_t w0 = v * 3; - v = w0 & 0x00FF00FF; - const uint32_t w1 = v * 3; - v = w1 & 0x00FF00FF; - row[30 + t / 2] = __vsub4(__byte_perm(w0, w1, 0x7531), 0x01010101); + for (int t = 0; t < 5; ++t) { + const uint32_t w_lo = v_lo * 3; + const uint32_t w_hi = v_hi * 3; + v_lo = w_lo & 0x00FF00FF; + v_hi = w_hi & 0x00FF00FF; + q[t] = __vsub4(__byte_perm(w_lo, w_hi, 0x7531), 0x01010101); + if (full_lane) { + row[dst_base + t * dst_stride] = q[t]; } } + if (lane == 6) { + // q[t] = {qh0.t, qh1.t, qh0.t, qh1.t}; the old layout is {qh0.t, qh1.t, qh0.t+1, qh1.t+1}. + row[30] = __byte_perm(q[0], q[1], 0x5410); + row[31] = __byte_perm(q[2], q[3], 0x5410); + } } constexpr int scale_entries_per_block = QK_PTQ1_0 / QK8_1;