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 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 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 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 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}