1#![no_std]
2
3#[cfg(feature = "std")]
4extern crate std;
5
6extern crate alloc;
7
8#[macro_use]
9extern crate derive_new;
10
11pub use cubecl_zspace as zspace;
12use cubecl_zspace::Shape;
13use cubecl_zspace::Strides;
14
15pub mod frontend;
17pub mod io;
19
20pub mod post_processing;
21
22pub use cubecl_common::future;
24
25use cubecl_ir::VectorSize;
26use cubecl_runtime::client::ComputeClient;
27pub use cubecl_runtime::memory_management::MemoryConfiguration;
28use cubecl_runtime::server::CubeCountSelection;
29pub use frontend::cmma;
30
31pub use cubecl_ir as ir;
33
34pub mod codegen;
35pub mod compute;
36pub mod prelude;
37
38mod pod;
39
40pub use codegen::*;
41pub use cubecl_runtime::runtime::*;
42pub use pod::*;
43
44pub use cubecl_macros::*;
45pub use cubecl_runtime::benchmark;
46pub use cubecl_runtime::client;
47pub use cubecl_runtime::compiler::{CompilationError, Compiler, CubeTask};
48pub use cubecl_runtime::memory_management::MemoryUsage;
49pub use cubecl_runtime::server;
50pub use cubecl_runtime::throughput;
51pub use cubecl_runtime::tune;
52
53use frontend::LaunchArg;
54
55pub use cubecl_common::*;
56
57pub use prelude::CubeCount;
58pub use prelude::{CubeDim, ExecutionMode};
59
60pub use num_traits;
61
62mod id;
63pub use id::*;
64
65#[doc(hidden)]
67pub mod __private {
68 pub use alloc::{format, vec};
69 pub use paste::paste;
70}
71
72pub use prelude::{Assign, IntoRuntime};
73
74pub fn calculate_cube_count_elemwise<R: Runtime>(
77 client: &ComputeClient<R>,
78 num_elems: usize,
79 cube_dim: CubeDim,
80) -> CubeCount {
81 if num_elems == 0 {
82 return CubeCount::Static(0, 0, 0);
83 }
84 let num_cubes = num_elems.div_ceil(cube_dim.num_elems() as usize);
85 CubeCountSelection::new(client, num_cubes as u32).cube_count()
86}
87
88pub fn tensor_vectorization_factor(
89 factors: &[VectorSize],
90 shape: &Shape,
91 strides: &Strides,
92 dim: usize,
93) -> VectorSize {
94 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
95}
96pub fn tensor_vectorization(
97 factors: &[VectorSize],
98 shape: &Shape,
99 strides: &Strides,
100 dim: usize,
101) -> VectorSize {
102 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
103}
104
105#[derive(Debug, Clone)]
106pub enum VectorizationError {
107 AxisOutOfBounds,
108 StrideMismatch,
109 NoValidVectorization,
110}
111
112pub fn tensor_vector_size_parallel(
125 optimized_vector_sizes: impl Iterator<Item = VectorSize>,
126 shape: &Shape,
127 strides: &Strides,
128 axis: usize,
129) -> VectorSize {
130 try_tensor_vector_size_parallel(optimized_vector_sizes, shape, strides, axis).unwrap_or(1)
131}
132
133pub fn try_tensor_vector_size_parallel(
135 supported_vector_sizes: impl Iterator<Item = VectorSize>,
136 shape: &Shape,
137 strides: &Strides,
138 axis: usize,
139) -> Result<VectorSize, VectorizationError> {
140 let stride = strides
141 .get(axis)
142 .ok_or(VectorizationError::AxisOutOfBounds)?;
143 if *stride != 1 {
144 return Err(VectorizationError::StrideMismatch);
145 }
146
147 let axis_shape = shape.get(axis).ok_or(VectorizationError::AxisOutOfBounds)?;
148
149 let next_stride = strides
156 .iter()
157 .enumerate()
158 .filter_map(|(i, &s)| (i != axis && s != 0).then_some(s))
159 .min()
160 .unwrap_or(0);
161
162 supported_vector_sizes
163 .filter(|&vector_size| axis_shape % vector_size == 0 && next_stride % vector_size == 0)
164 .max()
165 .ok_or(VectorizationError::NoValidVectorization)
166}
167
168pub fn tensor_vector_size_perpendicular(
179 supported_vector_sizes: impl Iterator<Item = VectorSize>,
180 shape: &[usize],
181 strides: &[usize],
182 axis: usize,
183) -> VectorSize {
184 try_tensor_vector_sizes_perpendicular(supported_vector_sizes, shape, strides, axis).unwrap_or(1)
185}
186
187pub fn try_tensor_vector_sizes_perpendicular(
189 supported_vector_sizes: impl Iterator<Item = VectorSize>,
190 shape: &[usize],
191 strides: &[usize],
192 axis: usize,
193) -> Result<VectorSize, VectorizationError> {
194 let axis_stride = strides
195 .get(axis)
196 .ok_or(VectorizationError::AxisOutOfBounds)?;
197
198 let prod_shape_axes_smaller_strides = strides
199 .iter()
200 .zip(shape.iter())
201 .filter(|(stride, _)| **stride < *axis_stride)
202 .map(|(_, shape)| shape)
203 .product::<usize>();
204
205 if *axis_stride != prod_shape_axes_smaller_strides {
206 return Err(VectorizationError::StrideMismatch);
207 }
208
209 supported_vector_sizes
210 .filter(|&vector_size| *axis_stride % vector_size == 0)
211 .max()
212 .ok_or(VectorizationError::NoValidVectorization)
213}
214
215pub type RuntimeArg<T, R> = <T as LaunchArg>::RuntimeArg<R>;
217pub type ExpandType<T> = <T as crate::prelude::CubeType>::ExpandType;
218
219#[cfg(feature = "export_tests")]
220pub mod runtime_tests;
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226
227 fn try_parallel(
228 sizes: &[VectorSize],
229 shape: &[usize],
230 strides: &[usize],
231 axis: usize,
232 ) -> Result<VectorSize, VectorizationError> {
233 try_tensor_vector_size_parallel(
234 sizes.iter().copied(),
235 &Shape::from(shape.iter().copied()),
236 &Strides::new(strides),
237 axis,
238 )
239 }
240
241 #[test]
242 fn parallel_contiguous_picks_max_vector_size() {
243 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 4, 1], 2).unwrap();
246 assert_eq!(v, 4);
247 }
248
249 #[test]
250 fn parallel_unfold_step_one_rejects_vectorization() {
251 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 1, 1], 2).unwrap();
257 assert_eq!(v, 1);
258 }
259
260 #[test]
261 fn parallel_unfold_step_two_allows_vectorization() {
262 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 2, 1], 2).unwrap();
266 assert_eq!(v, 2);
267 }
268
269 #[test]
270 fn parallel_broadcast_dim_ignored() {
271 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[0, 4, 1], 2).unwrap();
274 assert_eq!(v, 4);
275 }
276
277 #[test]
278 fn parallel_axis_stride_not_one_is_error() {
279 let err = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 1, 4], 2).unwrap_err();
280 assert!(matches!(err, VectorizationError::StrideMismatch));
281 }
282}