use cubecl::prelude::*;
#[allow(dead_code)]
pub const NF4_TABLE: [f32; 16] = [
-1.0, -0.6962, -0.5251, -0.3949, -0.2844, -0.1848, -0.0911, 0.0, 0.0796, 0.1609, 0.2461,
0.3379, 0.4407, 0.5626, 0.7230, 1.0,
];
#[allow(dead_code)]
pub const NF4_BOUNDARIES: [f32; 15] = [
-0.8481, -0.6107, -0.4600, -0.3397, -0.2346, -0.1380, -0.0456, 0.0398, 0.1203, 0.2035, 0.2920,
0.3893, 0.5017, 0.6428, 0.8615,
];
#[cube]
pub fn find_nearest_nf4<F: Float + CubeElement>(val: F) -> u32 {
if val < F::new(0.0) - F::new(0.8481) {
0u32.into()
} else if val < F::new(0.0) - F::new(0.6107) {
1u32.into()
} else if val < F::new(0.0) - F::new(0.4600) {
2u32.into()
} else if val < F::new(0.0) - F::new(0.3397) {
3u32.into()
} else if val < F::new(0.0) - F::new(0.2346) {
4u32.into()
} else if val < F::new(0.0) - F::new(0.1380) {
5u32.into()
} else if val < F::new(0.0) - F::new(0.0456) {
6u32.into()
} else if val < F::new(0.0398) {
7u32.into()
} else if val < F::new(0.1203) {
8u32.into()
} else if val < F::new(0.2035) {
9u32.into()
} else if val < F::new(0.2920) {
10u32.into()
} else if val < F::new(0.3893) {
11u32.into()
} else if val < F::new(0.5017) {
12u32.into()
} else if val < F::new(0.6428) {
13u32.into()
} else if val < F::new(0.8615) {
14u32.into()
} else {
15u32.into()
}
}
#[cube]
pub fn nf4_lookup<F: Float + CubeElement>(idx: u32) -> F {
if idx == 0u32 {
F::new(0.0) - F::new(1.0) } else if idx == 1u32 {
F::new(0.0) - F::new(0.6962) } else if idx == 2u32 {
F::new(0.0) - F::new(0.5251) } else if idx == 3u32 {
F::new(0.0) - F::new(0.3949) } else if idx == 4u32 {
F::new(0.0) - F::new(0.2844) } else if idx == 5u32 {
F::new(0.0) - F::new(0.1848) } else if idx == 6u32 {
F::new(0.0) - F::new(0.0911) } else if idx == 7u32 {
F::new(0.0)
} else if idx == 8u32 {
F::new(0.0796)
} else if idx == 9u32 {
F::new(0.1609)
} else if idx == 10u32 {
F::new(0.2461)
} else if idx == 11u32 {
F::new(0.3379)
} else if idx == 12u32 {
F::new(0.4407)
} else if idx == 13u32 {
F::new(0.5626)
} else if idx == 14u32 {
F::new(0.7230)
} else {
F::new(1.0)
}
}
#[cube(launch)]
pub fn nf4_quantize_kernel<F: Float + CubeElement>(
input: &Array<F>,
scales: &Array<F>,
output: &mut Array<u32>,
#[comptime] block_size: u32,
#[comptime] num_elements: u32,
) {
let out_idx = ABSOLUTE_POS;
let in_base = out_idx * 8usize;
if in_base >= (num_elements as usize) {
terminate!();
}
let scale_idx = in_base / (block_size as usize);
let scale = scales[scale_idx];
let inv_scale = if scale > F::new(1e-10) {
F::new(1.0) / scale
} else {
F::new(1.0)
};
let mut packed: u32 = 0u32;
let neg_one = F::new(0.0) - F::new(1.0);
#[unroll]
for i in 0usize..8usize {
let val_idx = in_base + i;
if val_idx < (num_elements as usize) {
let normalized = input[val_idx] * inv_scale;
let clamped = if normalized < neg_one {
neg_one
} else if normalized > F::new(1.0) {
F::new(1.0)
} else {
normalized
};
let nf4_idx = find_nearest_nf4::<F>(clamped);
packed = packed | (nf4_idx << ((i * 4usize) as u32));
}
}
output[out_idx] = packed;
}
#[cube(launch)]
pub fn nf4_dequantize_kernel<F: Float + CubeElement>(
input: &Array<u32>,
scales: &Array<F>,
output: &mut Array<F>,
#[comptime] block_size: u32,
#[comptime] num_elements: u32,
) {
let idx = ABSOLUTE_POS;
if idx >= (num_elements as usize) {
terminate!();
}
let packed_idx = idx / 8usize;
let sub_idx = idx % 8usize;
let packed = input[packed_idx];
let nf4_idx = (packed >> ((sub_idx * 4usize) as u32)) & 0xFu32;
let nf4_val: F = nf4_lookup::<F>(nf4_idx);
let scale_idx = idx / (block_size as usize);
let scale = scales[scale_idx];
output[idx] = nf4_val * scale;
}
#[cube(launch)]
pub fn compute_scales_kernel<F: Float + CubeElement>(
input: &Array<F>,
scales: &mut Array<F>,
#[comptime] block_size: u32,
#[comptime] num_blocks: u32,
) {
let block_idx = ABSOLUTE_POS;
if block_idx >= (num_blocks as usize) {
terminate!();
}
let start = block_idx * (block_size as usize);
let end = start + (block_size as usize);
let mut absmax = F::new(0.0);
for i in start..end {
let val = input[i];
let abs_val = if val < F::new(0.0) {
F::new(0.0) - val
} else {
val
};
if abs_val > absmax {
absmax = abs_val;
}
}
scales[block_idx] = if absmax > F::new(1e-10) {
absmax
} else {
F::new(1.0)
};
}
#[cube(launch)]
pub fn double_quantize_scales_kernel<F: Float + CubeElement>(
scales: &Array<F>,
quantized_scales: &mut Array<u32>,
scale_of_scales: &mut Array<F>,
#[comptime] num_blocks: u32,
#[comptime] blocks_per_superblock: u32,
) {
let superblock_idx = ABSOLUTE_POS;
let num_superblocks_u32 = (num_blocks + blocks_per_superblock - 1u32) / blocks_per_superblock;
if superblock_idx >= (num_superblocks_u32 as usize) {
terminate!();
}
let start = superblock_idx * (blocks_per_superblock as usize);
let mut absmax = F::new(0.0);
for i in 0u32..blocks_per_superblock {
let idx = start + (i as usize);
if idx < (num_blocks as usize) {
let val = scales[idx];
let abs_val = if val < F::new(0.0) {
F::new(0.0) - val
} else {
val
};
if abs_val > absmax {
absmax = abs_val;
}
}
}
let sos = if absmax > F::new(1e-10) {
absmax
} else {
F::new(1.0)
};
scale_of_scales[superblock_idx] = sos;
let inv_sos = F::new(127.0) / sos;
let packed_blocks_u32 = (blocks_per_superblock + 3u32) / 4u32;
let output_base = superblock_idx * (packed_blocks_u32 as usize);
for p in 0u32..packed_blocks_u32 {
let mut packed: u32 = 0u32;
#[unroll]
for j in 0u32..4u32 {
let idx = start + ((p * 4u32 + j) as usize);
if idx < (num_blocks as usize) {
let q_float = scales[idx] * inv_sos;
let neg_127 = F::new(0.0) - F::new(127.0);
let clamped = if q_float < neg_127 {
neg_127
} else if q_float > F::new(127.0) {
F::new(127.0)
} else {
q_float
};
let rounded = if clamped >= F::new(0.0) {
clamped + F::new(0.5)
} else {
clamped - F::new(0.5)
};
let q_i32 = i32::cast_from(rounded);
let q_u8 = (q_i32 + 127) as u32;
packed = packed | ((q_u8 & 0xFFu32) << (j * 8u32));
}
}
quantized_scales[(output_base + (p as usize))] = packed;
}
}
#[cube(launch)]
pub fn double_dequantize_scales_kernel<F: Float + CubeElement>(
quantized_scales: &Array<u32>,
scale_of_scales: &Array<F>,
scales: &mut Array<F>,
#[comptime] num_blocks: u32,
#[comptime] blocks_per_superblock: u32,
) {
let block_idx = ABSOLUTE_POS;
if block_idx >= (num_blocks as usize) {
terminate!();
}
let superblock_idx = block_idx / (blocks_per_superblock as usize);
let local_idx = block_idx % (blocks_per_superblock as usize);
let sos = scale_of_scales[superblock_idx];
let packed_per_superblock = ((blocks_per_superblock as usize) + 3usize) / 4usize;
let packed_idx = superblock_idx * packed_per_superblock + local_idx / 4usize;
let sub_idx = local_idx % 4usize;
let packed = quantized_scales[packed_idx];
let q_u8 = (packed >> ((sub_idx * 8usize) as u32)) & 0xFFu32;
let q_signed = (q_u8 as i32) - 127;
let dequantized = F::cast_from(q_signed) * sos / F::new(127.0);
scales[block_idx] = dequantized;
}
#[cube(launch)]
pub fn nf4_dequantize_double_quant_kernel<F: Float + CubeElement>(
input: &Array<u32>,
quantized_scales: &Array<u32>,
scale_of_scales: &Array<F>,
output: &mut Array<F>,
#[comptime] block_size: u32,
#[comptime] blocks_per_superblock: u32,
#[comptime] num_elements: u32,
) {
let idx = ABSOLUTE_POS;
if idx >= (num_elements as usize) {
terminate!();
}
let packed_idx = idx / 8usize;
let sub_idx = idx % 8usize;
let packed = input[packed_idx];
let nf4_idx = (packed >> ((sub_idx * 4usize) as u32)) & 0xFu32;
let nf4_val: F = nf4_lookup::<F>(nf4_idx);
let block_idx = idx / (block_size as usize);
let superblock_idx = block_idx / (blocks_per_superblock as usize);
let local_block_idx = block_idx % (blocks_per_superblock as usize);
let sos = scale_of_scales[superblock_idx];
let packed_per_superblock = ((blocks_per_superblock as usize) + 3usize) / 4usize;
let scale_packed_idx_u32 = superblock_idx * packed_per_superblock + local_block_idx / 4usize;
let scale_sub_idx_u32 = local_block_idx % 4usize;
let scale_packed = quantized_scales[scale_packed_idx_u32];
let q_u8 = (scale_packed >> ((scale_sub_idx_u32 * 8usize) as u32)) & 0xFFu32;
let q_signed = (q_u8 as i32) - 127;
let scale = F::cast_from(q_signed) * sos / F::new(127.0);
output[idx] = nf4_val * scale;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_nf4_boundaries() {
for i in 0..15 {
let midpoint = (NF4_TABLE[i] + NF4_TABLE[i + 1]) / 2.0;
assert!(
(NF4_BOUNDARIES[i] - midpoint).abs() < 0.001,
"Boundary {} mismatch: {} vs {}",
i,
NF4_BOUNDARIES[i],
midpoint
);
}
}
}