#define BLOCK_SIZE 32
#define BLOCKS_K ((HEAD_DIM_QK + BLOCK_SIZE - 1) / BLOCK_SIZE)
#define BLOCKS_V ((HEAD_DIM_V + BLOCK_SIZE - 1) / BLOCK_SIZE)

#if defined(K_Q4_0)
#define K_NQ 16
#define K_BLOCK_SIZE_BYTES 18u
#define K_BYTES_PER_THREAD 8u
#define K_BYTES_PER_INNER_LOOP 4u
#elif defined(K_Q8_0)
#define K_NQ 16
#define K_BLOCK_SIZE_BYTES 34u
#define K_BYTES_PER_THREAD 16u
#define K_BYTES_PER_INNER_LOOP 4u
#endif

#if defined(V_Q4_0)
#define V_NQ 16
#define V_BLOCK_SIZE_BYTES 18u
#define V_BYTES_PER_THREAD 8u
#define V_BYTES_PER_INNER_LOOP 4u
#elif defined(V_Q8_0)
#define V_NQ 16
#define V_BLOCK_SIZE_BYTES 34u
#define V_BYTES_PER_THREAD 16u
#define V_BYTES_PER_INNER_LOOP 4u
#endif

#if defined(K_Q4_0) || defined(K_Q8_0)
fn load_k_u16_at(byte_offset: u32) -> u32 {
    let word = K[byte_offset / 4u];
    let shift = (byte_offset & 2u) * 8u;
    return (word >> shift) & 0xFFFFu;
}

fn load_k_u32_at(byte_offset: u32) -> u32 {
    let word_idx = byte_offset / 4u;
    let shift = (byte_offset & 3u) * 8u;
    let lo = K[word_idx];
    if (shift == 0u) {
        return lo;
    }
    let hi = K[word_idx + 1u];
    return (lo >> shift) | (hi << (32u - shift));
}
#endif

#if defined(V_Q4_0) || defined(V_Q8_0)
fn load_v_u16_at(byte_offset: u32) -> u32 {
    let word = V[byte_offset / 4u];
    let shift = (byte_offset & 2u) * 8u;
    return (word >> shift) & 0xFFFFu;
}

fn load_v_u32_at(byte_offset: u32) -> u32 {
    let word_idx = byte_offset / 4u;
    let shift = (byte_offset & 3u) * 8u;
    let lo = V[word_idx];
    if (shift == 0u) {
        return lo;
    }
    let hi = V[word_idx + 1u];
    return (lo >> shift) | (hi << (32u - shift));
}
#endif

fn f16_from_u16(bits: u32) -> f16 {
    let packed = unpack2x16float(bits);
    return f16(packed[0]);
}

#if defined(K_Q4_0) || defined(K_Q8_0)
fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
    for (var elem_idx = local_x * K_NQ; elem_idx < kv_count * HEAD_DIM_QK; elem_idx += WG_SIZE * K_NQ) {
        let blck_idx = elem_idx / BLOCK_SIZE;
        let block_offset = (elem_idx % BLOCK_SIZE) / K_NQ;
        let k_row = blck_idx / BLOCKS_K;
        let global_k_row = kv_tile + k_row;
        let block_k = blck_idx % BLOCKS_K;
        let row_offset = k_row * HEAD_DIM_QK;
        let global_block_idx = k_head_offset + global_k_row * params.stride_k1 + block_k;
        let block_byte_base = global_block_idx * K_BLOCK_SIZE_BYTES;
        let d = f16_from_u16(load_k_u16_at(block_byte_base));
        let thread_byte_offset = block_offset * K_BYTES_PER_THREAD;
        let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
        for (var j = 0u; j < K_BYTES_PER_THREAD / K_BYTES_PER_INNER_LOOP; j += 1u) {
            let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * K_BYTES_PER_INNER_LOOP;
            let q_packed = load_k_u32_at(q_byte_offset);
#if defined(K_Q4_0)
            dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP);
#elif defined(K_Q8_0)
            dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP);
#endif
        }
    }
}
#endif

#if defined(V_Q4_0) || defined(V_Q8_0)
fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
    for (var elem_idx = local_x * V_NQ; elem_idx < kv_count * HEAD_DIM_V; elem_idx += WG_SIZE * V_NQ) {
        let blck_idx = elem_idx / BLOCK_SIZE;
        let block_offset = (elem_idx % BLOCK_SIZE) / V_NQ;
        let v_row = blck_idx / BLOCKS_V;
        let global_v_row = kv_tile + v_row;
        let block_k = blck_idx % BLOCKS_V;
        let row_offset = v_row * HEAD_DIM_V;
        let global_block_idx = v_head_offset + global_v_row * params.stride_v1 + block_k;
        let block_byte_base = global_block_idx * V_BLOCK_SIZE_BYTES;
        let d = f16_from_u16(load_v_u16_at(block_byte_base));
        let thread_byte_offset = block_offset * V_BYTES_PER_THREAD;
        let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
        for (var j = 0u; j < V_BYTES_PER_THREAD / V_BYTES_PER_INNER_LOOP; j += 1u) {
            let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * V_BYTES_PER_INNER_LOOP;
            let q_packed = load_v_u32_at(q_byte_offset);
#if defined(V_Q4_0)
            dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP);
#elif defined(V_Q8_0)
            dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP);
#endif
        }
    }
}
#endif
