spark_model/layers/ops/
q4k_mmq.rs1use anyhow::Result;
9use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
10use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
11
12pub const QK_K: u32 = 256;
14pub const Q4K_BLOCK_BYTES: usize = 144;
16pub const Q4K_MMQ_SMEM: u32 = 57856;
18const CUDA_QUANTIZE_BLOCK_SIZE_MMQ: u32 = 128;
19
20pub fn q4k_weight_bytes(nrows: u32, n_per_row: u32) -> usize {
22 (nrows as usize) * (n_per_row as usize / QK_K as usize) * Q4K_BLOCK_BYTES
23}
24
25pub fn q8_1_scratch_bytes(m: u32, k: u32) -> usize {
27 let kpad = div_ceil(k, QK_K) * QK_K;
28 (m as usize) * (kpad as usize) * 4 + (1 << 20)
29}
30
31pub fn dequant_nvfp4_to_bf16(
33 gpu: &dyn GpuBackend,
34 kernel: KernelHandle,
35 packed: DevicePtr,
36 scales: DevicePtr,
37 out_bf16: DevicePtr,
38 scale2: f32,
39 n: u32,
40 k: u32,
41 stream: u64,
42) -> Result<()> {
43 KernelLaunch::new(gpu, kernel)
44 .grid([n, 1, 1])
45 .block([256, 1, 1])
46 .arg_ptr(packed)
47 .arg_ptr(scales)
48 .arg_ptr(out_bf16)
49 .arg_f32(scale2)
50 .arg_u32(n)
51 .arg_u32(k)
52 .launch(stream)
53}
54
55pub fn quantize_weight_q4k(
57 gpu: &dyn GpuBackend,
58 kernel: KernelHandle,
59 input_bf16: DevicePtr,
60 out_q4k: DevicePtr,
61 nrows: u32,
62 n_per_row: u32,
63 stream: u64,
64) -> Result<()> {
65 let total_sb = (nrows as u64) * (n_per_row as u64 / QK_K as u64);
66 let grid_x = div_ceil(total_sb as u32, 128);
67 KernelLaunch::new(gpu, kernel)
68 .grid([grid_x, 1, 1])
69 .block([128, 1, 1])
70 .arg_ptr(input_bf16)
71 .arg_ptr(out_q4k)
72 .arg_u32(nrows)
73 .arg_u32(n_per_row)
74 .launch(stream)
75}
76
77pub fn quantize_act_q8_1(
79 gpu: &dyn GpuBackend,
80 kernel: KernelHandle, input_bf16: DevicePtr,
82 out_q8: DevicePtr,
83 m: u32,
84 k: u32,
85 stream: u64,
86) -> Result<()> {
87 let kpad = div_ceil(k, QK_K) * QK_K;
88 let grid_y = div_ceil(kpad, 4 * CUDA_QUANTIZE_BLOCK_SIZE_MMQ);
89 KernelLaunch::new(gpu, kernel)
90 .grid([m, grid_y, 1])
91 .block([CUDA_QUANTIZE_BLOCK_SIZE_MMQ, 1, 1])
92 .arg_ptr(input_bf16)
93 .arg_ptr(out_q8)
94 .arg_u64(k as u64) .arg_u64(k as u64) .arg_u64(kpad as u64) .arg_u32(m) .launch(stream)
99}
100
101pub fn q4k_mmq_gemm(
103 gpu: &dyn GpuBackend,
104 kernel_nc: KernelHandle, kernel_wc: KernelHandle, a_q8: DevicePtr, w_q4k: DevicePtr, out_bf16: DevicePtr,
109 m: u32,
110 n: u32,
111 k: u32,
112 stream: u64,
113) -> Result<()> {
114 let kernel = if !n.is_multiple_of(128) {
115 kernel_wc
116 } else {
117 kernel_nc
118 };
119 KernelLaunch::new(gpu, kernel)
120 .grid([div_ceil(n, 128), div_ceil(m, 128), 1])
121 .block([32, 8, 1])
122 .shared_mem(Q4K_MMQ_SMEM)
123 .arg_ptr(w_q4k) .arg_ptr(a_q8) .arg_ptr(out_bf16) .arg_u32(n) .arg_u32(m) .arg_u32(k) .arg_u32(k / QK_K) .arg_u32(m) .arg_u32(n) .launch(stream)
133}