diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp new file mode 100644 index 000000000000..44214a643caa --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp @@ -0,0 +1,160 @@ +#version 450 +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#include "mul_mat_vec_base.glsl" + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +#define K_PER_ITER 16 + +uint a_offset, b_offset, d_offset; + +void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, + const uint num_rows, const uint tid, const uint i) { + const uint col = i * BLOCK_SIZE + K_PER_ITER * tid; + const uint iqs = col & (QUANT_K - 1); // quant index, col is 16-aligned + const uint iybs = col - iqs; // y block start index + + uint ibs[NUM_ROWS]; + uint ibi = first_row * p.ncols; + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + ibs[n] = (ibi + col) / QUANT_K; // block index + ibi += p.ncols; + } + + // Phase 1: Issue A global loads first. A is the larger traffic source in + // token generation, and these loads are independent of B, so starting them + // earlier improves latency hiding before B loads and compute begin. + + // A quantized value loads + const uint ib32 = iqs / 32; + const uint iq = 16 * ib32 + (iqs & 15); + const uint qbase = iq / 4; + uint qs0[NUM_ROWS]; + uint qs1[NUM_ROWS]; + uint qs2[NUM_ROWS]; + uint qs3[NUM_ROWS]; + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const uint abase = a_offset + ibs[n]; + qs0[n] = data_a_packed32[abase].qs[qbase]; + qs1[n] = data_a_packed32[abase].qs[qbase + 1]; + qs2[n] = data_a_packed32[abase].qs[qbase + 2]; + qs3[n] = data_a_packed32[abase].qs[qbase + 3]; + } + + // A scale and master scale loads + float ds[NUM_ROWS]; + float dls[NUM_ROWS]; + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const uint abase = a_offset + ibs[n]; + const uint sl = (data_a[abase].scales_l[ib32 / 2] >> (4 * (ib32 & 1))) & 0xF; + const uint sh = (data_a[abase].scales_h >> (2 * ib32)) & 3; + dls[n] = float(int(sl | (sh << 4)) - 32); + ds[n] = float(data_a[abase].d); + } + + // Precompute per-row scale factors. + float scales[NUM_ROWS]; + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { scales[n] = ds[n] * dls[n]; } + + // B loads + vec4 bv0[NUM_COLS]; + vec4 bv1[NUM_COLS]; + vec4 bv2[NUM_COLS]; + vec4 bv3[NUM_COLS]; + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + const uint bbase = j * p.batch_stride_b + b_offset + iybs + iqs; + const uint bidx = bbase / 4; + bv0[j] = vec4(data_b_v4[bidx]); + bv1[j] = vec4(data_b_v4[bidx + 1]); + bv2[j] = vec4(data_b_v4[bidx + 2]); + bv3[j] = vec4(data_b_v4[bidx + 3]); + } + + // Phase 2: Interleave LDS LUT dequantization with dot-product compute. + // Dequantizing each vec4 immediately before its dot() use spreads LDS + // traffic across the compute phase, reduces peak live dequantized values, + // and allows LDS latency to overlap with v_dot4_f32 execution. + const uint qshift = (iqs & 16) >> 2; + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + const vec4 bv0j = bv0[j]; + const vec4 bv1j = bv1[j]; + const vec4 bv2j = bv2[j]; + const vec4 bv3j = bv3[j]; + + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const u8vec4 us0 = unpack8((qs0[n] >> qshift) & 0x0F0F0F0F); + vec4 vs0 = vec4(kvalues_iq4nl[us0.x], kvalues_iq4nl[us0.y], + kvalues_iq4nl[us0.z], kvalues_iq4nl[us0.w]); + const u8vec4 us1 = unpack8((qs1[n] >> qshift) & 0x0F0F0F0F); + vec4 vs1 = vec4(kvalues_iq4nl[us1.x], kvalues_iq4nl[us1.y], + kvalues_iq4nl[us1.z], kvalues_iq4nl[us1.w]); + const u8vec4 us2 = unpack8((qs2[n] >> qshift) & 0x0F0F0F0F); + vec4 vs2 = vec4(kvalues_iq4nl[us2.x], kvalues_iq4nl[us2.y], + kvalues_iq4nl[us2.z], kvalues_iq4nl[us2.w]); + const u8vec4 us3 = unpack8((qs3[n] >> qshift) & 0x0F0F0F0F); + vec4 vs3 = vec4(kvalues_iq4nl[us3.x], kvalues_iq4nl[us3.y], + kvalues_iq4nl[us3.z], kvalues_iq4nl[us3.w]); + + FLOAT_TYPE rowtmp = dot(bv0j, vs0) + dot(bv1j, vs1) + dot(bv2j, vs2) + dot(bv3j, vs3); + rowtmp *= scales[n]; + temp[j][n] += rowtmp; + } + } +} + +void compute_outputs(const uint32_t first_row, const uint32_t num_rows) { + const uint tid = gl_LocalInvocationID.x; + + get_offsets(a_offset, b_offset, d_offset); + + FLOAT_TYPE temp[NUM_COLS][NUM_ROWS]; + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) { + temp[j][i] = FLOAT_TYPE(0); + } + } + + uint num_iters = p.ncols / (K_PER_ITER * BLOCK_SIZE); + if (num_iters * K_PER_ITER * BLOCK_SIZE + K_PER_ITER * tid < p.ncols) { + num_iters++; + } + + // Single uniform unrolling strategy: fully unroll small K, two-way unroll + // large K. + if (num_iters <= 8) { + [[unroll]] for (uint i = 0; i < num_iters; ++i) { + iter(temp, first_row, num_rows, tid, i * K_PER_ITER); + } + } else { + uint i = 0; + const uint unrolled = num_iters & ~1u; + while (i < unrolled) { + iter(temp, first_row, num_rows, tid, i * K_PER_ITER); + iter(temp, first_row, num_rows, tid, (i + 1) * K_PER_ITER); + i += 2; + } + if (i < num_iters) { + iter(temp, first_row, num_rows, tid, i * K_PER_ITER); + } + } + + reduce_result(temp, d_offset, first_row, num_rows, tid); +} + +void main() { + const uint first_row = NUM_ROWS * (gl_WorkGroupID.x + gl_NumWorkGroups.x * gl_WorkGroupID.z); + + init_iq_shmem(gl_WorkGroupSize); + + // do NUM_ROWS at a time, unless there aren't enough remaining rows + if (first_row + NUM_ROWS <= p.stride_d) { + compute_outputs(first_row, NUM_ROWS); + } else { + if (first_row >= p.stride_d) { + return; + } + compute_outputs(first_row, min(NUM_ROWS, p.stride_d - first_row)); + } +} \ No newline at end of file diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 27ff68c10d5b..845a601cff66 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -734,7 +734,7 @@ void process_shaders() { for (const auto& tname : type_names) { // mul mat vec std::string data_a_key = "DATA_A_" + to_uppercase(tname); - std::string shader = (string_ends_with(tname, "_k") || string_starts_with(tname, "iq1_") || string_starts_with(tname, "iq2_") || string_starts_with(tname, "iq3_") || tname == "tq2_0") ? "mul_mat_vec_" + tname + ".comp" : "mul_mat_vec.comp"; + std::string shader = (string_ends_with(tname, "_k") || string_starts_with(tname, "iq1_") || string_starts_with(tname, "iq2_") || string_starts_with(tname, "iq3_") || tname == "iq4_xs" || tname == "tq2_0") ? "mul_mat_vec_" + tname + ".comp" : "mul_mat_vec.comp"; string_to_spv("mul_mat_vec_" + tname + "_f32_f32", shader, merge_maps(base_dict, {{data_a_key, "1"}, {"B_TYPE", "float"}, {"B_TYPEV2", "vec2"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}})); string_to_spv("mul_mat_vec_" + tname + "_f16_f32", shader, merge_maps(base_dict, {{data_a_key, "1"}, {"B_TYPE", "float16_t"}, {"B_TYPEV2", "f16vec2"}, {"B_TYPEV4", "f16vec4"}, {"D_TYPE", "float"}}));