Skip to content
Closed
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
160 changes: 160 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp
Original file line number Diff line number Diff line change
@@ -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));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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"}}));
Expand Down