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
23fn 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
33fn 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
44pub 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
53pub 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 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 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 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 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}