Skip to main content

burn_cubecl/ops/
qtensor.rs

1use burn_backend::{
2    Bytes, DType, ExecutionError, Shape, SplitPolicy, TensorData, TensorMetadata, TensorPrimitive,
3    get_device_settings,
4    ops::QTensorOps,
5    quantization::{
6        QParamTensor, QuantLevel, QuantMode, QuantParam, QuantPropagation, QuantScheme, QuantValue,
7        QuantizationParametersPrimitive, params_shape,
8    },
9    tensor::{Device, FloatTensor, QuantizedTensor},
10};
11use burn_std::{FloatDType, Metadata};
12use cubecl::server::{MemoryLayout, MemoryLayoutDescriptor, MemoryLayoutStrategy};
13use cubecl::{e2m1x2, quant::scheme::QuantStore};
14
15use crate::{
16    CubeBackend, CubeRuntime,
17    kernel::{self, matmul::MatmulStrategy},
18    tensor::{CubeTensor, QParams},
19};
20
21use super::{into_data, permute, swap_dims};
22
23/// Create a quantized tensor with packed values (u32).
24fn new_qtensor_optimized<R: CubeRuntime>(
25    data: Bytes,
26    shape: impl Into<Shape>,
27    scheme: QuantScheme,
28    device: &R::Device,
29) -> CubeTensor<R> {
30    new_qtensor(data, shape, scheme, device, MemoryLayoutStrategy::Optimized)
31}
32
33/// Create a quantized tensor with packed values (u32).
34fn new_qtensor<R: CubeRuntime>(
35    data: Bytes,
36    shape: impl Into<Shape>,
37    scheme: QuantScheme,
38    device: &R::Device,
39    kind: MemoryLayoutStrategy,
40) -> CubeTensor<R> {
41    new_quantized(shape, scheme, device, Some(data), kind)
42}
43
44/// Create an empty quantized tensor.
45pub fn empty_qtensor_optimized<R: CubeRuntime>(
46    shape: impl Into<Shape>,
47    scheme: QuantScheme,
48    device: &R::Device,
49) -> CubeTensor<R> {
50    empty_qtensor(shape, scheme, device, MemoryLayoutStrategy::Optimized)
51}
52
53/// Create an empty quantized tensor.
54pub fn empty_qtensor<R: CubeRuntime>(
55    shape: impl Into<Shape>,
56    scheme: QuantScheme,
57    device: &R::Device,
58    kind: MemoryLayoutStrategy,
59) -> CubeTensor<R> {
60    new_quantized(shape, scheme, device, None, kind)
61}
62
63fn new_quantized<R: CubeRuntime>(
64    shape: impl Into<Shape>,
65    scheme: QuantScheme,
66    device: &R::Device,
67    data: Option<Bytes>,
68    alloc_kind: MemoryLayoutStrategy,
69) -> CubeTensor<R> {
70    let client = R::client(device);
71    let shape: Shape = shape.into();
72    let mut shape_value: Shape = shape.clone();
73
74    let rank = shape.rank();
75    let shape_last = shape[rank - 1];
76    let num_quants = scheme.num_quants();
77
78    let data_size = match scheme.store {
79        QuantStore::PackedU32(_) => {
80            if !shape_last.is_multiple_of(num_quants) {
81                panic!("Can't store in u32")
82            }
83            shape_value[rank - 1] = shape_last.div_ceil(num_quants);
84            size_of::<u32>()
85        }
86        QuantStore::Native => match scheme.value {
87            QuantValue::Q8F | QuantValue::Q8S | QuantValue::E4M3 | QuantValue::E5M2 => {
88                size_of::<i8>()
89            }
90            QuantValue::Q4F
91            | QuantValue::Q4S
92            | QuantValue::Q2F
93            | QuantValue::Q2S
94            | QuantValue::E2M1 => {
95                panic!("Can't store native sub-byte values")
96            }
97        },
98        QuantStore::PackedNative(_) => match scheme.value {
99            QuantValue::E2M1 => size_of::<e2m1x2>(),
100            other => panic!("{other:?} doesn't support native packing"),
101        },
102    };
103
104    let scales_dtype = match scheme.param {
105        QuantParam::F32 => DType::F32,
106        QuantParam::F16 => DType::F16,
107        QuantParam::BF16 => DType::BF16,
108        // Represented by U8 and reinterpreted in the kernel
109        QuantParam::UE8M0 | QuantParam::UE4M3 => DType::U8,
110    };
111
112    let scales_shape = params_shape(&shape, scheme.level);
113    let data_desc = MemoryLayoutDescriptor::new(alloc_kind, shape_value.clone(), data_size);
114    let scales_desc =
115        MemoryLayoutDescriptor::new(alloc_kind, scales_shape.clone(), scales_dtype.size());
116
117    let mut tensors = match data {
118        Some(data) => {
119            let num_bytes = shape_value.num_elements() * data_size;
120
121            match data.split(num_bytes, SplitPolicy::Shared) {
122                Ok((bytes_data, bytes_scales)) => client
123                    .create_tensors(vec![(data_desc, bytes_data), (scales_desc, bytes_scales)]),
124                Err((data, _)) => client.create_tensors_from_slices(vec![
125                    (data_desc, &data[..num_bytes]),
126                    (scales_desc, &data[num_bytes..]),
127                ]),
128            }
129        }
130        None => client.empty_tensors(vec![data_desc, scales_desc]),
131    };
132    let MemoryLayout {
133        memory: scales_handle,
134        strides: scales_strides,
135    } = tensors.remove(1);
136    let MemoryLayout { memory, strides } = tensors.remove(0);
137
138    let scales = QParamTensor {
139        offset_start: scales_handle.offset_start.unwrap_or(0) as usize,
140        offset_end: scales_handle.offset_end.unwrap_or(0) as usize,
141        metadata: Metadata::new(scales_shape, scales_strides),
142        dtype: scales_dtype,
143    };
144    let qparams = QParams { scales };
145
146    CubeTensor::new_quantized(
147        client,
148        memory,
149        shape,
150        device.clone(),
151        strides,
152        DType::QFloat(scheme),
153        qparams,
154    )
155}
156
157impl<R: CubeRuntime> QTensorOps<Self> for CubeBackend<R> {
158    fn q_from_data(data: TensorData, device: &Device<Self>) -> QuantizedTensor<Self> {
159        match data.dtype {
160            DType::QFloat(scheme) => match scheme {
161                QuantScheme {
162                    level: QuantLevel::BlockTensor { .. },
163                    ..
164                } => unimplemented!("two-level quantization is not supported yet"),
165                QuantScheme {
166                    level: QuantLevel::Tensor | QuantLevel::Block(_),
167                    mode: QuantMode::Symmetric,
168                    value:
169                        QuantValue::Q8F
170                        | QuantValue::Q8S
171                        | QuantValue::Q4F
172                        | QuantValue::Q4S
173                        | QuantValue::Q2F
174                        | QuantValue::Q2S
175                        | QuantValue::E4M3
176                        | QuantValue::E5M2
177                        | QuantValue::E2M1,
178                    ..
179                } => {
180                    // TensorData quantized representation is the same, with multiple quantized values
181                    // packed into u32 and quantization parameters appended to the bytes
182                    new_qtensor_optimized(data.bytes, data.shape.clone(), scheme, device)
183                }
184            },
185            _ => panic!(
186                "Invalid dtype (expected DType::QFloat, got {:?})",
187                data.dtype
188            ),
189        }
190    }
191
192    // TODO: quantize_dynamic (we can compute min-max on the fly and scale, especially when not per-tensor)
193
194    fn quantize(
195        tensor: FloatTensor<Self>,
196        scheme: &QuantScheme,
197        qparams: QuantizationParametersPrimitive<Self>,
198    ) -> QuantizedTensor<Self> {
199        kernel::quantization::quantize(tensor, scheme, qparams.scales)
200    }
201
202    fn dequantize(tensor: QuantizedTensor<Self>, dtype: FloatDType) -> FloatTensor<Self> {
203        kernel::quantization::dequantize(tensor, dtype.into())
204    }
205
206    fn q_to_device(tensor: QuantizedTensor<Self>, device: &Device<Self>) -> QuantizedTensor<Self> {
207        super::to_device(tensor, device)
208    }
209
210    fn q_reshape(tensor: QuantizedTensor<Self>, shape: Shape) -> QuantizedTensor<Self> {
211        super::q_reshape(tensor, shape)
212    }
213
214    async fn q_into_data(tensor: QuantizedTensor<Self>) -> Result<TensorData, ExecutionError> {
215        if tensor.qparams.is_none() {
216            return into_data(tensor).await;
217        }
218
219        let (shape, dtype) = (tensor.shape(), tensor.dtype);
220        let (values, params) = tensor.quantized_handles().unwrap();
221
222        let mut data_values = into_data(values).await?;
223        let data_params = into_data(params).await?;
224
225        data_values.bytes.extend_from_byte_slice(&data_params.bytes);
226
227        Ok(TensorData {
228            bytes: data_values.bytes,
229            shape,
230            dtype,
231        })
232    }
233
234    fn q_swap_dims(
235        tensor: QuantizedTensor<Self>,
236        dim1: usize,
237        dim2: usize,
238    ) -> QuantizedTensor<Self> {
239        swap_dims(tensor, dim1, dim2)
240    }
241
242    fn q_permute(tensor: QuantizedTensor<Self>, axes: &[usize]) -> QuantizedTensor<Self> {
243        permute(tensor, axes)
244    }
245
246    fn q_flip(_tensor: QuantizedTensor<Self>, _axes: &[usize]) -> QuantizedTensor<Self> {
247        unimplemented!()
248    }
249
250    fn q_matmul(lhs: TensorPrimitive<Self>, rhs: TensorPrimitive<Self>) -> TensorPrimitive<Self> {
251        let (settings, scheme) = match (&lhs, &rhs) {
252            (TensorPrimitive::QFloat(lhs), _) => {
253                (get_device_settings::<Self>(&lhs.device), lhs.scheme())
254            }
255            (_, TensorPrimitive::QFloat(rhs)) => {
256                (get_device_settings::<Self>(&rhs.device), rhs.scheme())
257            }
258            _ => unreachable!(),
259        };
260
261        // Inherit precision for mixed inputs, default to `FloatElem` for fully quantized.
262        let out_dtype = match (&lhs, &rhs) {
263            (TensorPrimitive::Float(lhs), _) => lhs.dtype,
264            (_, TensorPrimitive::Float(rhs)) => rhs.dtype,
265            _ => settings.float_dtype.into(),
266        };
267
268        let (_lhs_dtype, lhs) = match lhs {
269            TensorPrimitive::Float(lhs) => (lhs.dtype, lhs),
270            TensorPrimitive::QFloat(lhs) => (out_dtype, lhs),
271        };
272        let (_rhs_dtype, rhs) = match rhs {
273            TensorPrimitive::Float(rhs) => (rhs.dtype, rhs),
274            TensorPrimitive::QFloat(rhs) => (out_dtype, rhs),
275        };
276
277        let out =
278            kernel::matmul::matmul(lhs, rhs, None, MatmulStrategy::default(), out_dtype).unwrap();
279
280        match settings.quantization.propagation {
281            QuantPropagation::Propagate => {
282                TensorPrimitive::QFloat(Self::quantize_dynamic(out, &scheme))
283            }
284            QuantPropagation::Inhibit => TensorPrimitive::Float(out),
285        }
286    }
287}