spark_model/layers/ops/
gemm_fp4.rs1#![allow(unused_imports)]
7
8use anyhow::Result;
9use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
10use spark_runtime::kernel_args::{KernelLaunch, div_ceil};
11
12use crate::weight_map::QuantizedWeight;
13
14use super::*;
15
16#[allow(clippy::too_many_arguments)]
20pub fn quantize_bf16_to_nvfp4(
21 gpu: &dyn GpuBackend,
22 kernel: KernelHandle,
23 input: DevicePtr,
24 packed_out: DevicePtr,
25 scale_out: DevicePtr,
26 m: u32,
27 k: u32,
28 stream: u64,
29) -> Result<()> {
30 KernelLaunch::new(gpu, kernel)
31 .grid([m, 1, 1])
32 .block([128, 1, 1])
33 .arg_ptr(input)
34 .arg_ptr(packed_out)
35 .arg_ptr(scale_out)
36 .arg_f32(1.0) .arg_u32(m) .arg_u32(k)
39 .launch(stream)
40}
41
42#[allow(clippy::too_many_arguments)]
47pub fn w4a4_gemm(
48 gpu: &dyn GpuBackend,
49 kernel: KernelHandle,
50 a_packed: DevicePtr,
51 a_scale: DevicePtr,
52 weight: &QuantizedWeight,
53 output: DevicePtr,
54 m: u32,
55 n: u32,
56 k: u32,
57 stream: u64,
58) -> Result<()> {
59 KernelLaunch::new(gpu, kernel)
60 .grid([div_ceil(n, 128), div_ceil(m, 128), 1])
61 .block([256, 1, 1])
62 .arg_ptr(a_packed)
63 .arg_ptr(a_scale)
64 .arg_ptr(weight.weight)
65 .arg_ptr(weight.weight_scale)
66 .arg_ptr(output)
67 .arg_f32(1.0) .arg_f32(weight.weight_scale_2) .arg_u32(m)
70 .arg_u32(n)
71 .arg_u32(k)
72 .launch(stream)
73}
74
75#[allow(clippy::too_many_arguments)]
79pub fn w4a4_gemm_mfast(
82 gpu: &dyn GpuBackend,
83 kernel: KernelHandle,
84 a_packed: DevicePtr,
85 a_scale: DevicePtr,
86 weight: &QuantizedWeight,
87 output: DevicePtr,
88 m: u32,
89 n: u32,
90 k: u32,
91 stream: u64,
92) -> Result<()> {
93 KernelLaunch::new(gpu, kernel)
94 .grid([div_ceil(m, 128), div_ceil(n, 128), 1])
95 .block([128, 1, 1])
96 .arg_ptr(a_packed)
97 .arg_ptr(a_scale)
98 .arg_ptr(weight.weight)
99 .arg_ptr(weight.weight_scale)
100 .arg_ptr(output)
101 .arg_f32(1.0) .arg_f32(weight.weight_scale_2) .arg_u32(m)
104 .arg_u32(n)
105 .arg_u32(k)
106 .launch(stream)
107}