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