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::memory_management::{
51 InstallMemoryPoolsError, MemoryPoolKind, MemoryPoolReport, MemoryReport,
52};
53pub use cubecl_runtime::server;
54pub use cubecl_runtime::throughput;
55pub use cubecl_runtime::tune;
56
57use frontend::LaunchArg;
58
59pub use cubecl_common::*;
60
61pub use prelude::CubeCount;
62pub use prelude::{CubeDim, ExecutionMode};
63
64pub use num_traits;
65
66mod id;
67pub use id::*;
68
69#[doc(hidden)]
71pub mod __private {
72 pub use alloc::{format, vec};
73 pub use paste::paste;
74}
75
76pub use prelude::{Assign, IntoRuntime};
77
78pub fn calculate_cube_count_elemwise<R: Runtime>(
81 client: &ComputeClient<R>,
82 num_elems: usize,
83 cube_dim: CubeDim,
84) -> CubeCount {
85 if num_elems == 0 {
86 return CubeCount::Static(0, 0, 0);
87 }
88 let num_cubes = num_elems.div_ceil(cube_dim.num_elems() as usize);
89 CubeCountSelection::new(client, num_cubes as u32).cube_count()
90}
91
92pub fn tensor_vectorization_factor(
93 factors: &[VectorSize],
94 shape: &Shape,
95 strides: &Strides,
96 dim: usize,
97) -> VectorSize {
98 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
99}
100pub fn tensor_vectorization(
101 factors: &[VectorSize],
102 shape: &Shape,
103 strides: &Strides,
104 dim: usize,
105) -> VectorSize {
106 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
107}
108
109#[derive(Debug, Clone)]
110pub enum VectorizationError {
111 AxisOutOfBounds,
112 StrideMismatch,
113 NoValidVectorization,
114}
115
116pub fn tensor_vector_size_parallel(
129 optimized_vector_sizes: impl Iterator<Item = VectorSize>,
130 shape: &Shape,
131 strides: &Strides,
132 axis: usize,
133) -> VectorSize {
134 try_tensor_vector_size_parallel(optimized_vector_sizes, shape, strides, axis).unwrap_or(1)
135}
136
137pub fn try_tensor_vector_size_parallel(
139 supported_vector_sizes: impl Iterator<Item = VectorSize>,
140 shape: &Shape,
141 strides: &Strides,
142 axis: usize,
143) -> Result<VectorSize, VectorizationError> {
144 let stride = strides
145 .get(axis)
146 .ok_or(VectorizationError::AxisOutOfBounds)?;
147 if *stride != 1 {
148 return Err(VectorizationError::StrideMismatch);
149 }
150
151 let axis_shape = shape.get(axis).ok_or(VectorizationError::AxisOutOfBounds)?;
152
153 let next_stride = strides
160 .iter()
161 .enumerate()
162 .filter_map(|(i, &s)| (i != axis && s != 0).then_some(s))
163 .min()
164 .unwrap_or(0);
165
166 supported_vector_sizes
167 .filter(|&vector_size| axis_shape % vector_size == 0 && next_stride % vector_size == 0)
168 .max()
169 .ok_or(VectorizationError::NoValidVectorization)
170}
171
172pub fn tensor_vector_size_perpendicular(
183 supported_vector_sizes: impl Iterator<Item = VectorSize>,
184 shape: &[usize],
185 strides: &[usize],
186 axis: usize,
187) -> VectorSize {
188 try_tensor_vector_sizes_perpendicular(supported_vector_sizes, shape, strides, axis).unwrap_or(1)
189}
190
191pub fn try_tensor_vector_sizes_perpendicular(
193 supported_vector_sizes: impl Iterator<Item = VectorSize>,
194 shape: &[usize],
195 strides: &[usize],
196 axis: usize,
197) -> Result<VectorSize, VectorizationError> {
198 let axis_stride = strides
199 .get(axis)
200 .ok_or(VectorizationError::AxisOutOfBounds)?;
201
202 let prod_shape_axes_smaller_strides = strides
203 .iter()
204 .zip(shape.iter())
205 .filter(|(stride, _)| **stride < *axis_stride)
206 .map(|(_, shape)| shape)
207 .product::<usize>();
208
209 if *axis_stride != prod_shape_axes_smaller_strides {
210 return Err(VectorizationError::StrideMismatch);
211 }
212
213 supported_vector_sizes
214 .filter(|&vector_size| *axis_stride % vector_size == 0)
215 .max()
216 .ok_or(VectorizationError::NoValidVectorization)
217}
218
219pub type RuntimeArg<T, R> = <T as LaunchArg>::RuntimeArg<R>;
221pub type ExpandType<T> = <T as crate::prelude::CubeType>::ExpandType;
222
223#[cfg(feature = "export_tests")]
224pub mod runtime_tests;
226
227#[cfg(test)]
228mod tests {
229 use super::*;
230
231 fn try_parallel(
232 sizes: &[VectorSize],
233 shape: &[usize],
234 strides: &[usize],
235 axis: usize,
236 ) -> Result<VectorSize, VectorizationError> {
237 try_tensor_vector_size_parallel(
238 sizes.iter().copied(),
239 &Shape::from(shape.iter().copied()),
240 &Strides::new(strides),
241 axis,
242 )
243 }
244
245 #[test]
246 fn parallel_contiguous_picks_max_vector_size() {
247 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 4, 1], 2).unwrap();
250 assert_eq!(v, 4);
251 }
252
253 #[test]
254 fn parallel_unfold_step_one_rejects_vectorization() {
255 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 1, 1], 2).unwrap();
261 assert_eq!(v, 1);
262 }
263
264 #[test]
265 fn parallel_unfold_step_two_allows_vectorization() {
266 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 2, 1], 2).unwrap();
270 assert_eq!(v, 2);
271 }
272
273 #[test]
274 fn parallel_broadcast_dim_ignored() {
275 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[0, 4, 1], 2).unwrap();
278 assert_eq!(v, 4);
279 }
280
281 #[test]
282 fn parallel_axis_stride_not_one_is_error() {
283 let err = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 1, 4], 2).unwrap_err();
284 assert!(matches!(err, VectorizationError::StrideMismatch));
285 }
286}