|
|
|
@@ -100,34 +100,37 @@ const BLOCK_SIZE_BYTES = 18u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
override BLOCKS_K = TILE_K/BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
|
|
|
|
|
const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < F16_PER_THREAD; j += 2) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 2u + 2u * (block_offset + j);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
|
|
|
|
|
let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte(q_packed, k);
|
|
|
|
|
let q_hi = (f16((q_byte >> 4) & 0xF) - 8.0) * d;
|
|
|
|
|
let q_lo = (f16(q_byte & 0xF) - 8.0) * d;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -141,35 +144,38 @@ const BLOCK_SIZE_BYTES = 20u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
override BLOCKS_K = TILE_K/BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
|
|
|
|
|
const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
let m = load_f16_at_src0(block_byte_base + 2u);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < F16_PER_THREAD; j += 2) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 4u + 2u * (block_offset + j);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
|
|
|
|
|
let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte(q_packed, k);
|
|
|
|
|
let q_lo = f16(q_byte & 0xF) * d + m;
|
|
|
|
|
let q_hi = f16((q_byte >> 4) & 0xF) * d + m;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -178,52 +184,49 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
|
|
#endif // INIT_SRC0_SHMEM_Q4_1
|
|
|
|
|
|
|
|
|
|
#ifdef INIT_SRC0_SHMEM_Q5_0
|
|
|
|
|
// 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
|
|
|
|
|
const BLOCK_SIZE = 32u;
|
|
|
|
|
const BLOCK_SIZE_BYTES = 22u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
// tile_k is defined as 32u, so blocks_k ends up being 1 always
|
|
|
|
|
override BLOCKS_K = TILE_K / BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
|
|
|
|
|
const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
let qh_packed = load_u32_at_src0(block_byte_base + 2u);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < 2; j++) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 6u + 2u * (block_offset + j * 2u);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 6u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
|
|
|
|
|
let j_adjusted = j + (block_offset / 2u);
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte(q_packed, k);
|
|
|
|
|
|
|
|
|
|
let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
|
|
|
|
|
let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
|
|
|
|
|
let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
|
|
|
|
|
let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi) - 16.0) * d;
|
|
|
|
|
let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
|
|
|
|
|
let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
|
|
|
|
|
let q_lo = (f16((q_byte & 0xF) | qh_lo) - 16.0) * d;
|
|
|
|
|
|
|
|
|
|
shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
|
|
|
|
|
shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -232,54 +235,49 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
|
|
#endif // INIT_SRC0_SHMEM_Q5_0
|
|
|
|
|
|
|
|
|
|
#ifdef INIT_SRC0_SHMEM_Q5_1
|
|
|
|
|
// 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
|
|
|
|
|
const BLOCK_SIZE = 32u;
|
|
|
|
|
const BLOCK_SIZE_BYTES = 24u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
// tile_k is defined as 32u, so blocks_k ends up being 1 always
|
|
|
|
|
override BLOCKS_K = TILE_K / BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
|
|
|
|
|
const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
let m = load_f16_at_src0(block_byte_base + 2u);
|
|
|
|
|
let qh_packed = load_u32_at_src0(block_byte_base + 4u);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < 2; j++) {
|
|
|
|
|
|
|
|
|
|
let q_byte_offset = block_byte_base + 8u + 2u * (block_offset + j * 2u);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 8u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
|
|
|
|
|
let j_adjusted = j + (block_offset / 2u);
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte(q_packed, k);
|
|
|
|
|
|
|
|
|
|
let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
|
|
|
|
|
let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi)) * d + m;
|
|
|
|
|
let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
|
|
|
|
|
let q_lo = (f16((q_byte & 0xF) | qh_lo)) * d + m;
|
|
|
|
|
|
|
|
|
|
shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
|
|
|
|
|
shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
|
|
|
|
|
let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
|
|
|
|
|
let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
|
|
|
|
|
let q_hi = f16(((q_byte >> 4) & 0xF) | qh_hi) * d + m;
|
|
|
|
|
let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
|
|
|
|
|
let q_lo = f16((q_byte & 0xF) | qh_lo) * d + m;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -293,33 +291,34 @@ const BLOCK_SIZE_BYTES = 34u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
override BLOCKS_K = TILE_K/BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread
|
|
|
|
|
const BYTES_PER_THREAD = 16u; // NQ(16) weights use 16 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < F16_PER_THREAD; j+=2) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 2u + 2u * (block_offset + j);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte_i32(q_packed, k);
|
|
|
|
|
|
|
|
|
|
let q_val = f16(q_byte) * d;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k] = q_val;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -333,34 +332,35 @@ const BLOCK_SIZE_BYTES = 36u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
override BLOCKS_K = TILE_K/BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
|
|
|
|
|
const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread, 2 threads per block
|
|
|
|
|
const BYTES_PER_THREAD = 16u; // NQ(16) weights use 16 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
let d = load_f16_at_src0(block_byte_base);
|
|
|
|
|
let m = load_f16_at_src0(block_byte_base + 2u);
|
|
|
|
|
|
|
|
|
|
for (var j = 0u; j < F16_PER_THREAD; j+=2) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 4u + 2u * (block_offset + j);
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
for (var k = 0u; k < 4u; k++) {
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte_i32(q_packed, k);
|
|
|
|
|
|
|
|
|
|
let q_val = f16(q_byte) * d + m;
|
|
|
|
|
shmem[shmem_idx + j * 2 + k] = q_val;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
@@ -1163,3 +1163,48 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
#endif // INIT_SRC0_SHMEM_IQ3_S
|
|
|
|
|
|
|
|
|
|
#ifdef INIT_SRC0_SHMEM_MXFP4
|
|
|
|
|
const BLOCK_SIZE = 32u;
|
|
|
|
|
const BLOCK_SIZE_BYTES = 17u;
|
|
|
|
|
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
|
|
|
|
|
override BLOCKS_K = TILE_K/BLOCK_SIZE;
|
|
|
|
|
const NQ = 16u;
|
|
|
|
|
const BYTES_PER_THREAD = 8u; // NQ(16) weights uses 8 bytes of q
|
|
|
|
|
const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
|
|
|
|
|
|
|
|
|
|
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
|
|
|
|
|
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
|
|
|
|
|
let blck_idx = i / BLOCK_SIZE;
|
|
|
|
|
let block_offset = (i % BLOCK_SIZE) / NQ;
|
|
|
|
|
let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
|
|
|
|
|
|
|
|
|
|
let tile_m = blck_idx / BLOCKS_K;
|
|
|
|
|
let global_m = offset_m + tile_m;
|
|
|
|
|
let block_k = blck_idx % BLOCKS_K;
|
|
|
|
|
let global_block_k = k_outer / BLOCK_SIZE + block_k;
|
|
|
|
|
|
|
|
|
|
if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
|
|
|
|
|
let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
|
|
|
|
|
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
|
|
|
|
|
let eu8 = get_byte(load_u32_at_src0(block_byte_base), 0);
|
|
|
|
|
let e = ldexp(1.0, i32(eu8) - 128);
|
|
|
|
|
|
|
|
|
|
// store NQ(16) weights
|
|
|
|
|
for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
|
|
|
|
|
|
|
|
|
|
let q_byte_offset = block_byte_base + 1u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
|
|
|
|
|
let q_packed = load_u32_at_src0(q_byte_offset);
|
|
|
|
|
|
|
|
|
|
for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
|
|
|
|
|
let q_byte = get_byte(q_packed, k);
|
|
|
|
|
let q_hi = f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * e;
|
|
|
|
|
let q_lo = f32(kvalues_mxfp4[q_byte & 0xF]) * e;
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = f16(q_lo);
|
|
|
|
|
shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = f16(q_hi);
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
#endif // INIT_SRC0_SHMEM_MXFP4
|
|
|
|
|