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