diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl index d17ed72f2498..bcdce6a7a4ed 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl @@ -592,15 +592,15 @@ float rocmfpx_fp3_dequant(uint ib, uint idx, uint a_offset) { return float(rocmfpx_fp3_decode_code(rocmfpx_fp3_get_bits(ib, idx * 3u, a_offset))) * d; } +// idx is a multiple of 4 at every call site, so bit_pos is a multiple of 12 and +// the shift is 0 or 4: the twelve bits always fit in two bytes, no third byte +// and no branch. int32_t rocmfpx_fp3_pack4_window(uint ib, uint idx, uint a_offset) { const uint bit_pos = idx * 3u; const uint byte_pos = bit_pos >> 3u; const uint sh = bit_pos & 7u; uint bits = uint(data_a[a_offset + ib].qs[byte_pos]) | (uint(data_a[a_offset + ib].qs[byte_pos + 1u]) << 8); - if (sh > 4u) { - bits |= uint(data_a[a_offset + ib].qs[byte_pos + 2u]) << 16; - } bits = (bits >> sh) & 0xFFFu; return pack32(i8vec4(kvalues_rocmfpx_fp3_const[ bits & 7u], kvalues_rocmfpx_fp3_const[(bits >> 3) & 7u], @@ -610,10 +610,9 @@ int32_t rocmfpx_fp3_pack4_window(uint ib, uint idx, uint a_offset) { vec4 rocmfpx_fp3_dequant4(uint ib, uint idx, uint a_offset) { const vec4 q = vec4(unpack8(rocmfpx_fp3_pack4_window(ib, idx, a_offset))); - return q * vec4(ue4m3_to_fp32(data_a[a_offset + ib].e[(idx + 0u) >= 16u ? 1u : 0u]), - ue4m3_to_fp32(data_a[a_offset + ib].e[(idx + 1u) >= 16u ? 1u : 0u]), - ue4m3_to_fp32(data_a[a_offset + ib].e[(idx + 2u) >= 16u ? 1u : 0u]), - ue4m3_to_fp32(data_a[a_offset + ib].e[(idx + 3u) >= 16u ? 1u : 0u])); + // idx is a multiple of 4 and the halves split at element 16, so all four + // weights share one scale + return q * ue4m3_to_fp32(data_a[a_offset + ib].e[idx >= 16u ? 1u : 0u]); } vec2 dequantize(uint ib, uint iqs, uint a_offset) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp index 068eb40a7a9a..6ed8fe34267b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp @@ -17,7 +17,12 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; #if defined(DATA_A_Q2_0) || defined(DATA_A_QUANT_K) #define K_PER_ITER 16 -#elif defined(DATA_A_QUANT_LEGACY) || defined(DATA_A_MXFP4) || defined(DATA_A_ROCMFP4) || defined(DATA_A_ROCMFP4_FAST) || defined(DATA_A_ROCMFPX_FP6) || defined(DATA_A_ROCMFPX_FP8) +#elif defined(DATA_A_ROCMFP4) || defined(DATA_A_ROCMFP4_FAST) +// A whole block per call. The ROCmFP4 scales are UE4M3 and cost a byte load plus +// a table lookup each, so decoding them once per 32 weights instead of once per +// 8 is worth the extra B registers. +#define K_PER_ITER 32 +#elif defined(DATA_A_QUANT_LEGACY) || defined(DATA_A_MXFP4) || defined(DATA_A_ROCMFPX_FP6) || defined(DATA_A_ROCMFPX_FP8) #define K_PER_ITER 8 #elif defined(DATA_A_IQ1_S) || defined(DATA_A_IQ1_M) || defined(DATA_A_ROCMFPX_FP2) || defined(DATA_A_ROCMFPX_FP3) #define K_PER_ITER 32 @@ -46,9 +51,16 @@ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const cache_b_ds = vec2(data_b[b_block_idx_outer].ds[b_block_idx_inner]); #if QUANT_R == 2 - // Assumes K_PER_ITER == 8 +#if K_PER_ITER == 8 cache_b_qs[0] = data_b[b_block_idx_outer].qs[b_block_idx_inner * 8 + b_qs_idx]; cache_b_qs[1] = data_b[b_block_idx_outer].qs[b_block_idx_inner * 8 + b_qs_idx + 4]; +#elif K_PER_ITER == 32 + [[unroll]] for (uint k = 0; k < 8; ++k) { + cache_b_qs[k] = data_b[b_block_idx_outer].qs[b_block_idx_inner * 8 + k]; + } +#else +#error unimplemented +#endif #else #if K_PER_ITER == 8 cache_b_qs[0] = data_b[b_block_idx_outer].qs[b_block_idx_inner * 8 + b_qs_idx * 2]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl index 95f431d70af6..94f57c65acc0 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl @@ -175,21 +175,33 @@ FLOAT_TYPE mul_q8_1(const int32_t q_sum, const float da, const vec2 dsb, const i #endif #if defined(DATA_A_ROCMFP4_FAST) +// K_PER_ITER is 32 for this type: one call consumes a whole block, so the +// single UE4M3 scale is decoded once per block rather than once per 8 weights. FLOAT_TYPE mmvq_dot_product(const uint ib_a, const uint iqs) { - const i32vec2 data_a_qs = repack(ib_a, iqs); + int32_t q_sum = 0; - const int32_t q_sum = dotPacked4x8EXT(data_a_qs.x, cache_b_qs[0]) + - dotPacked4x8EXT(data_a_qs.y, cache_b_qs[1]); + [[unroll]] for (uint k = 0; k < 4; ++k) { + const i32vec2 data_a_qs = repack(ib_a, k); + q_sum += dotPacked4x8EXT(data_a_qs.x, cache_b_qs[k]) + + dotPacked4x8EXT(data_a_qs.y, cache_b_qs[k + 4]); + } const FLOAT_TYPE d = FLOAT_TYPE(ue4m3_to_fp32(data_a[ib_a].e)); return FLOAT_TYPE(cache_b_ds.x * float(q_sum) * d); } #elif defined(DATA_A_ROCMFP4) +// K_PER_ITER is 32 for this type: one call consumes a whole block, so the +// two UE4M3 scales are decoded once per block rather than once per 8 weights. FLOAT_TYPE mmvq_dot_product(const uint ib_a, const uint iqs) { - const i32vec2 data_a_qs = repack(ib_a, iqs); - - const int32_t q_sum0 = dotPacked4x8EXT(data_a_qs.x, cache_b_qs[0]); - const int32_t q_sum1 = dotPacked4x8EXT(data_a_qs.y, cache_b_qs[1]); + // the two halves carry different scales, so they cannot share an accumulator + int32_t q_sum0 = 0; + int32_t q_sum1 = 0; + + [[unroll]] for (uint k = 0; k < 4; ++k) { + const i32vec2 data_a_qs = repack(ib_a, k); + q_sum0 += dotPacked4x8EXT(data_a_qs.x, cache_b_qs[k]); + q_sum1 += dotPacked4x8EXT(data_a_qs.y, cache_b_qs[k + 4]); + } const FLOAT_TYPEV2 d = get_dm(ib_a); return FLOAT_TYPE(cache_b_ds.x * (float(q_sum0) * d.x + float(q_sum1) * d.y));