Skip to main content

cubek_std/
input_binding.rs

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