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::Client;
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 pod::*;
43
44pub use cubecl_macros::*;
45pub use cubecl_runtime::benchmark;
46pub use cubecl_runtime::client;
47pub use cubecl_runtime::compiler::{CompilationError, Compiler};
48pub use cubecl_runtime::kernel::{CubeKernel, PrecompiledSource};
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::sync::Arc;
73 pub use alloc::{format, vec};
74 pub use cubecl_runtime::runtime::Runtime;
75 pub use paste::paste;
76}
77
78pub use prelude::{Assign, IntoRuntime};
79
80pub fn calculate_cube_count_elemwise(
83 client: &Client,
84 num_elems: usize,
85 cube_dim: CubeDim,
86) -> CubeCount {
87 if num_elems == 0 {
88 return CubeCount::Static(0, 0, 0);
89 }
90 let num_cubes = num_elems.div_ceil(cube_dim.num_elems() as usize);
91 CubeCountSelection::new(client, num_cubes as u32).cube_count()
92}
93
94pub fn tensor_vectorization_factor(
95 factors: &[VectorSize],
96 shape: &Shape,
97 strides: &Strides,
98 dim: usize,
99) -> VectorSize {
100 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
101}
102pub fn tensor_vectorization(
103 factors: &[VectorSize],
104 shape: &Shape,
105 strides: &Strides,
106 dim: usize,
107) -> VectorSize {
108 tensor_vector_size_parallel(factors.iter().cloned(), shape, strides, dim)
109}
110
111#[derive(Debug, Clone)]
112pub enum VectorizationError {
113 AxisOutOfBounds,
114 StrideMismatch,
115 NoValidVectorization,
116}
117
118pub fn tensor_vector_size_parallel(
131 optimized_vector_sizes: impl Iterator<Item = VectorSize>,
132 shape: &Shape,
133 strides: &Strides,
134 axis: usize,
135) -> VectorSize {
136 try_tensor_vector_size_parallel(optimized_vector_sizes, shape, strides, axis).unwrap_or(1)
137}
138
139pub fn try_tensor_vector_size_parallel(
141 supported_vector_sizes: impl Iterator<Item = VectorSize>,
142 shape: &Shape,
143 strides: &Strides,
144 axis: usize,
145) -> Result<VectorSize, VectorizationError> {
146 let stride = strides
147 .get(axis)
148 .ok_or(VectorizationError::AxisOutOfBounds)?;
149 if *stride != 1 {
150 return Err(VectorizationError::StrideMismatch);
151 }
152
153 let axis_shape = shape.get(axis).ok_or(VectorizationError::AxisOutOfBounds)?;
154
155 let next_stride = strides
162 .iter()
163 .enumerate()
164 .filter_map(|(i, &s)| (i != axis && s != 0).then_some(s))
165 .min()
166 .unwrap_or(0);
167
168 supported_vector_sizes
169 .filter(|&vector_size| axis_shape % vector_size == 0 && next_stride % vector_size == 0)
170 .max()
171 .ok_or(VectorizationError::NoValidVectorization)
172}
173
174pub fn tensor_vector_size_perpendicular(
185 supported_vector_sizes: impl Iterator<Item = VectorSize>,
186 shape: &[usize],
187 strides: &[usize],
188 axis: usize,
189) -> VectorSize {
190 try_tensor_vector_sizes_perpendicular(supported_vector_sizes, shape, strides, axis).unwrap_or(1)
191}
192
193pub fn try_tensor_vector_sizes_perpendicular(
195 supported_vector_sizes: impl Iterator<Item = VectorSize>,
196 shape: &[usize],
197 strides: &[usize],
198 axis: usize,
199) -> Result<VectorSize, VectorizationError> {
200 let axis_stride = strides
201 .get(axis)
202 .ok_or(VectorizationError::AxisOutOfBounds)?;
203
204 let prod_shape_axes_smaller_strides = strides
205 .iter()
206 .zip(shape.iter())
207 .filter(|(stride, _)| **stride < *axis_stride)
208 .map(|(_, shape)| shape)
209 .product::<usize>();
210
211 if *axis_stride != prod_shape_axes_smaller_strides {
212 return Err(VectorizationError::StrideMismatch);
213 }
214
215 supported_vector_sizes
216 .filter(|&vector_size| *axis_stride % vector_size == 0)
217 .max()
218 .ok_or(VectorizationError::NoValidVectorization)
219}
220
221pub type RuntimeArg<T> = <T as LaunchArg>::RuntimeArg;
223pub type ExpandType<T> = <T as crate::prelude::CubeType>::ExpandType;
224
225#[cfg(feature = "export_tests")]
226pub mod runtime_tests;
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232
233 fn try_parallel(
234 sizes: &[VectorSize],
235 shape: &[usize],
236 strides: &[usize],
237 axis: usize,
238 ) -> Result<VectorSize, VectorizationError> {
239 try_tensor_vector_size_parallel(
240 sizes.iter().copied(),
241 &Shape::from(shape.iter().copied()),
242 &Strides::new(strides),
243 axis,
244 )
245 }
246
247 #[test]
248 fn parallel_contiguous_picks_max_vector_size() {
249 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 4, 1], 2).unwrap();
252 assert_eq!(v, 4);
253 }
254
255 #[test]
256 fn parallel_unfold_step_one_rejects_vectorization() {
257 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 1, 1], 2).unwrap();
263 assert_eq!(v, 1);
264 }
265
266 #[test]
267 fn parallel_unfold_step_two_allows_vectorization() {
268 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[12, 2, 1], 2).unwrap();
272 assert_eq!(v, 2);
273 }
274
275 #[test]
276 fn parallel_broadcast_dim_ignored() {
277 let v = try_parallel(&[1, 2, 4], &[1, 9, 4], &[0, 4, 1], 2).unwrap();
280 assert_eq!(v, 4);
281 }
282
283 #[test]
284 fn parallel_axis_stride_not_one_is_error() {
285 let err = try_parallel(&[1, 2, 4], &[1, 9, 4], &[36, 1, 4], 2).unwrap_err();
286 assert!(matches!(err, VectorizationError::StrideMismatch));
287 }
288}