Skip to main content

cubek_std/
input_binding.rs

1use cubecl::std::tensor::{into_contiguous_packed, into_contiguous_pitched};
2use cubecl::{
3    client::Client,
4    frontend::Scalar,
5    ir::{AddressType, ElemType},
6    prelude::TensorBinding,
7    server::LaunchError,
8    zspace::Shape,
9};
10use cubecl_common::quant::scheme::{QuantScheme, QuantStore, QuantValue};
11
12#[derive(Debug)]
13#[allow(clippy::large_enum_variant)]
14pub enum InputBinding {
15    Normal(TensorBinding, ElemType),
16    Quantized {
17        data: TensorBinding,
18        data_dtype: ElemType,
19        scale: TensorBinding,
20        scale_dtype: ElemType,
21        /// Unpacked shape, excluding padding
22        shape: Shape,
23        scheme: QuantScheme,
24    },
25}
26
27impl Clone for InputBinding {
28    fn clone(&self) -> Self {
29        match self {
30            Self::Normal(arg0, arg1) => Self::Normal(arg0.clone(), *arg1),
31            Self::Quantized {
32                data,
33                data_dtype,
34                scale,
35                scale_dtype,
36                shape,
37                scheme,
38            } => Self::Quantized {
39                data: data.clone(),
40                data_dtype: *data_dtype,
41                scale: scale.clone(),
42                scale_dtype: *scale_dtype,
43                shape: shape.clone(),
44                scheme: *scheme,
45            },
46        }
47    }
48}
49
50impl InputBinding {
51    pub fn new(data: TensorBinding, dtype: ElemType) -> Self {
52        Self::Normal(data, dtype)
53    }
54
55    pub fn swap_dims(&mut self, dim0: usize, dim1: usize) {
56        match self {
57            Self::Normal(handle, _dtype) => {
58                handle.shape.swap(dim0, dim1);
59                handle.strides.swap(dim0, dim1);
60            }
61            Self::Quantized {
62                data,
63                scale,
64                shape,
65                scheme,
66                data_dtype: _,
67                scale_dtype: _,
68            } => {
69                if scheme.num_levels() > 1 {
70                    unimplemented!("two-level quantization is not supported here, got {scheme:?}");
71                }
72
73                let rank = data.shape.len();
74
75                data.shape.swap(dim0, dim1);
76                data.strides.swap(dim0, dim1);
77
78                // Swap dims for scale and block size if block scaled quant is used
79                if scheme.block_size().is_some() {
80                    scale.shape.swap(dim0, dim1);
81                    scale.strides.swap(dim0, dim1);
82                    scheme.swap_block_dims(rank, dim0, dim1);
83                }
84
85                shape.swap(dim0, dim1);
86
87                // Swap packed dim if packed dim is either of `dim0` or `dim1`
88                if let QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim) =
89                    &mut scheme.store
90                {
91                    if *packed_dim == rank - dim0 - 1 {
92                        *packed_dim = rank - dim1 - 1;
93                    } else if *packed_dim == rank - dim1 - 1 {
94                        *packed_dim = rank - dim0 - 1;
95                    }
96                }
97            }
98        }
99    }
100    pub fn quantized(
101        data: TensorBinding,
102        scale: TensorBinding,
103        shape: Shape,
104        scheme: QuantScheme,
105        data_dtype: ElemType,
106        scale_dtype: ElemType,
107    ) -> Self {
108        Self::Quantized {
109            data,
110            scale,
111            shape,
112            scheme,
113            data_dtype,
114            scale_dtype,
115        }
116    }
117
118    pub fn data(&self) -> &TensorBinding {
119        match self {
120            InputBinding::Normal(handle, ..) => handle,
121            InputBinding::Quantized { data, .. } => data,
122        }
123    }
124
125    pub fn data_elem_size(&self) -> usize {
126        match self {
127            InputBinding::Normal(_, ty) => ty.size(),
128            InputBinding::Quantized { data_dtype, .. } => data_dtype.size(),
129        }
130    }
131
132    pub fn into_data(self) -> TensorBinding {
133        match self {
134            InputBinding::Normal(handle, ..) => handle,
135            InputBinding::Quantized { data, .. } => data,
136        }
137    }
138
139    pub fn data_mut(&mut self) -> &mut TensorBinding {
140        match self {
141            InputBinding::Normal(handle, ..) => handle,
142            InputBinding::Quantized { data, .. } => data,
143        }
144    }
145
146    pub fn scale(&self) -> Option<&TensorBinding> {
147        match self {
148            InputBinding::Normal(..) => None,
149            InputBinding::Quantized { scale, .. } => Some(scale),
150        }
151    }
152
153    pub fn scheme(&self) -> Option<&QuantScheme> {
154        match self {
155            InputBinding::Normal(..) => None,
156            InputBinding::Quantized { scheme, .. } => Some(scheme),
157        }
158    }
159
160    pub fn shape(&self) -> &Shape {
161        match self {
162            InputBinding::Normal(handle, ..) => &handle.shape,
163            InputBinding::Quantized { shape, .. } => shape,
164        }
165    }
166
167    pub fn into_contiguous(self, client: &Client) -> Result<Self, LaunchError> {
168        let val = match self {
169            Self::Normal(data, dtype) => Self::Normal(
170                into_contiguous_pitched(client, data, dtype).binding(),
171                dtype,
172            ),
173            Self::Quantized {
174                data,
175                scale,
176                shape,
177                scheme,
178                data_dtype,
179                scale_dtype,
180            } => {
181                let mut scheme = scheme;
182                let data = match scheme.store {
183                    // e2m1 has native packing (e2m1x2) so also needs to be re-packed
184                    QuantStore::PackedNative(packed_dim) if scheme.value == QuantValue::E2M1 => {
185                        let mut data = into_contiguous_packed(
186                            client,
187                            data,
188                            packed_dim,
189                            &shape,
190                            scheme.num_quants(),
191                            u8::elem_type_native(),
192                        );
193                        scheme = scheme.with_store(QuantStore::PackedNative(0));
194                        data.dtype = data_dtype;
195                        data
196                    }
197                    QuantStore::PackedU32(packed_dim) => {
198                        let mut data = into_contiguous_packed(
199                            client,
200                            data,
201                            packed_dim,
202                            &shape,
203                            scheme.num_quants(),
204                            u32::elem_type_native(),
205                        );
206                        data.dtype = data_dtype;
207                        scheme = scheme.with_store(QuantStore::PackedU32(0));
208                        data
209                    }
210                    _ => into_contiguous_pitched(client, data, data_dtype),
211                };
212
213                Self::Quantized {
214                    data: data.binding(),
215                    scale,
216                    shape,
217                    scheme,
218                    data_dtype,
219                    scale_dtype,
220                }
221            }
222        };
223
224        Ok(val)
225    }
226
227    pub fn required_address_type(&self) -> AddressType {
228        match self {
229            InputBinding::Normal(handle, ty) => handle.required_address_type(ty.size()),
230            InputBinding::Quantized {
231                data,
232                shape,
233                scheme,
234                ..
235            } => {
236                let handle_addr = data.required_address_type(scheme.size_bits_stored() / 8);
237                let conceptual_addr = AddressType::from_len(shape.iter().product());
238                handle_addr.max(conceptual_addr)
239            }
240        }
241    }
242}