Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -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],
Expand All @@ -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) {
Expand Down
16 changes: 14 additions & 2 deletions ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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];
Expand Down
26 changes: 19 additions & 7 deletions ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl
Original file line number Diff line number Diff line change
Expand Up @@ -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));
Expand Down