#version 450
#extension GL_EXT_shader_explicit_arithmetic_types : require
#include "mul_mat_vec_base.glsl"
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
FLOAT_TYPE temp[NUM_COLS][NUM_ROWS];
// Walks the packed bytes directly (byte m, digit t) rather than via
// tq1_0_byte_of()/tq1_0_digit_of(): one byte per thread, expanded in place.
void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
uint a_offset, b_offset, d_offset;
get_offsets(a_offset, b_offset, d_offset);
const uint num_blocks_per_row = p.ncols / QUANT_K;
const uint tid = gl_LocalInvocationID.x;
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
[[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) {
temp[j][i] = FLOAT_TYPE(0);
}
}
for (uint nrow = 0; nrow < num_rows; ++nrow) {
const uint ib0 = a_offset + (first_row + nrow) * num_blocks_per_row;
for (uint jcol = 0; jcol < NUM_COLS; ++jcol) {
const uint b_base = (jcol * p.batch_stride_b);
for (uint i = tid/8; i < num_blocks_per_row; i += gl_WorkGroupSize.x/8) {
const FLOAT_TYPE d = float(data_a[ib0 + i].d);
// First qs chunk: 32 bytes (5*32 elements)
[[unroll]] for (uint m = tid%8; m < 32; m += 8) {
const uint q_byte = uint(data_a[ib0 + i].qs[m]);
[[unroll]] for (uint t = 0; t < 5; ++t) {
const uint xi = tq1_0_trit(q_byte, t);
const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
const uint elem = t * 32u + m;
const uint b_idx = i * QUANT_K + elem;
temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
}
}
// Second qs chunk: 16 bytes (5*16 elements)
[[unroll]] for (uint m = tid%8; m < 16; m += 8) {
const uint q_byte = uint(data_a[ib0 + i].qs[32u + m]);
[[unroll]] for (uint t = 0; t < 5; ++t) {
const uint xi = tq1_0_trit(q_byte, t);
const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
const uint elem = 160u + t * 16u + m;
const uint b_idx = i * QUANT_K + elem;
temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
}
}
// qh bytes: 4 bytes (4*4 elements)
[[unroll]] for (uint j = tid%8; j < 4; j += 8) {
const uint qh_byte = uint(data_a[ib0 + i].qh[j]);
[[unroll]] for (uint t = 0; t < 4; ++t) {
const uint xi = tq1_0_trit(qh_byte, t);
const FLOAT_TYPE dequant_val = FLOAT_TYPE(d * (float(xi) - 1.0f));
const uint elem = 240u + t * 4u + j;
const uint b_idx = i * QUANT_K + elem;
temp[jcol][nrow] += dequant_val * FLOAT_TYPE(data_b[b_base + b_offset + b_idx]);
}
}
}
}
}
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);
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, p.stride_d - first_row);
}
}