use crate::mul_mm::{f16_to_f32, MulMmKind, SUB};
pub const Q8_0: MulMmKind = MulMmKind {
name: "Q8_0",
module_name: "ferrox_mul_mm_q8_0",
fn_name: "q8_0_mul_mm",
block_bytes: 34,
block_elems: 32,
dequant_src: r#"
__device__ __forceinline__ void ferrox_dequant_sub(
const unsigned char* xb, int il, float* reg
) {
const float d = ferrox_f16_to_f32(
(unsigned short)xb[0] | ((unsigned short)xb[1] << 8));
const signed char* qs = (const signed char*)(xb + 2) + 16 * il;
#pragma unroll
for (int i = 0; i < 16; i++) {
reg[i] = (float)qs[i] * d;
}
}
"#,
dequant_twin: dequant_sub_q8_0,
codebook: None,
};
fn dequant_sub_q8_0(xb: &[u8], il: usize, reg: &mut [f32; SUB]) {
let d = f16_to_f32(u16::from(xb[0]) | (u16::from(xb[1]) << 8));
let qs = &xb[2 + SUB * il..2 + SUB * il + SUB];
for (r, q) in reg.iter_mut().zip(qs.iter()) {
*r = f32::from(*q as i8) * d;
}
}
pub const Q4_0: MulMmKind = MulMmKind {
name: "Q4_0",
module_name: "ferrox_mul_mm_q4_0",
fn_name: "q4_0_mul_mm",
block_bytes: 18,
block_elems: 32,
dequant_src: r#"
__device__ __forceinline__ void ferrox_dequant_sub(
const unsigned char* xb, int il, float* reg
) {
const float d = ferrox_f16_to_f32(
(unsigned short)xb[0] | ((unsigned short)xb[1] << 8));
const unsigned char* qs = xb + 2;
const float d1 = il ? d / 16.0f : d;
const float d2 = d1 / 256.0f;
const float md = -8.0f * d;
const unsigned short mask0 = il ? 0x00F0 : 0x000F;
const unsigned short mask1 = (unsigned short)(mask0 << 8);
#pragma unroll
for (int i = 0; i < 8; i++) {
const unsigned short w =
(unsigned short)qs[2 * i] | ((unsigned short)qs[2 * i + 1] << 8);
reg[2 * i + 0] = d1 * (float)(w & mask0) + md;
reg[2 * i + 1] = d2 * (float)(w & mask1) + md;
}
}
"#,
dequant_twin: dequant_sub_q4_0,
codebook: None,
};
fn dequant_sub_q4_0(xb: &[u8], il: usize, reg: &mut [f32; SUB]) {
let d = f16_to_f32(u16::from(xb[0]) | (u16::from(xb[1]) << 8));
let qs = &xb[2..2 + 16];
let d1 = if il != 0 { d / 16.0 } else { d };
let d2 = d1 / 256.0;
let md = -8.0 * d;
let mask0: u16 = if il != 0 { 0x00F0 } else { 0x000F };
let mask1: u16 = mask0 << 8;
for i in 0..8 {
let w = u16::from(qs[2 * i]) | (u16::from(qs[2 * i + 1]) << 8);
reg[2 * i] = d1 * f32::from(w & mask0) + md;
reg[2 * i + 1] = d2 * f32::from(w & mask1) + md;
}
}
pub const Q5_0: MulMmKind = MulMmKind {
name: "Q5_0",
module_name: "ferrox_mul_mm_q5_0",
fn_name: "q5_0_mul_mm",
block_bytes: 22,
block_elems: 32,
dequant_src: r#"
__device__ __forceinline__ void ferrox_dequant_sub(
const unsigned char* xb, int il, float* reg
) {
const float d = ferrox_f16_to_f32(
(unsigned short)xb[0] | ((unsigned short)xb[1] << 8));
const float md = -16.0f * d;
const unsigned int qh = (unsigned int)xb[2]
| ((unsigned int)xb[3] << 8)
| ((unsigned int)xb[4] << 16)
| ((unsigned int)xb[5] << 24);
const unsigned char* qs = xb + 6;
const unsigned short mask = il ? 0x00F0 : 0x000F;
const int x_mv = il ? 4 : 0;
const int gh_mv = il ? 12 : 0;
const int gh_bk = il ? 0 : 4;
#pragma unroll
for (int i = 0; i < 8; i++) {
const unsigned short w =
(unsigned short)qs[2 * i] | ((unsigned short)qs[2 * i + 1] << 8);
const unsigned char xh_0 =
(unsigned char)(((qh >> (gh_mv + 2 * i)) << gh_bk) & 0x10u);
const unsigned char xh_1 =
(unsigned char)(((qh >> (gh_mv + 2 * i + 1)) << gh_bk) & 0x10u);
const int x0 = (int)((((w) & mask) >> x_mv) | xh_0);
const int x1 = (int)((((w >> 8) & mask) >> x_mv) | xh_1);
reg[2 * i + 0] = d * (float)x0 + md;
reg[2 * i + 1] = d * (float)x1 + md;
}
}
"#,
dequant_twin: dequant_sub_q5_0,
codebook: None,
};
fn dequant_sub_q5_0(xb: &[u8], il: usize, reg: &mut [f32; SUB]) {
let d = f16_to_f32(u16::from(xb[0]) | (u16::from(xb[1]) << 8));
let md = -16.0 * d;
let qh = u32::from(xb[2])
| (u32::from(xb[3]) << 8)
| (u32::from(xb[4]) << 16)
| (u32::from(xb[5]) << 24);
let qs = &xb[6..6 + 16];
let mask: u16 = if il != 0 { 0x00F0 } else { 0x000F };
let x_mv = if il != 0 { 4 } else { 0 };
let gh_mv = if il != 0 { 12 } else { 0 };
let gh_bk = if il != 0 { 0 } else { 4 };
for i in 0..8 {
let w = u16::from(qs[2 * i]) | (u16::from(qs[2 * i + 1]) << 8);
let xh_0 = (((qh >> (gh_mv + 2 * i)) << gh_bk) & 0x10) as u8;
let xh_1 = (((qh >> (gh_mv + 2 * i + 1)) << gh_bk) & 0x10) as u8;
let x0 = i32::from((((w & mask) >> x_mv) as u8) | xh_0);
let x1 = i32::from(((((w >> 8) & mask) >> x_mv) as u8) | xh_1);
reg[2 * i] = d * x0 as f32 + md;
reg[2 * i + 1] = d * x1 as f32 + md;
}
}