Skip to main content

cubecl_std/throughput/runners/
memory_direct.rs

1use cubecl::prelude::*;
2use cubecl_core as cubecl;
3use cubecl_runtime::throughput::{KernelConfig, ThroughputKey};
4
5use crate::throughput::LaunchConfig;
6
7/// Per-buffer size, clamped to the device's maximum allocation.
8const TARGET_BYTES: usize = 512 * 1024 * 1024;
9
10/// Builds the copy kernel.
11pub fn build_kernel<R: Runtime>(
12    client: &ComputeClient<R>,
13    key: ThroughputKey,
14    config: LaunchConfig,
15) -> KernelConfig {
16    let client = client.clone();
17    let dtype = key.dtype();
18
19    let line_bytes = config.vector_size * dtype.size();
20
21    let max_alloc = client.properties().memory.max_page_size as usize;
22    let target = TARGET_BYTES.min(max_alloc);
23
24    let total_threads = config.cube_count * config.cube_dim;
25    let num_lines = (target / line_bytes).max(total_threads);
26    let bytes = num_lines * line_bytes;
27
28    let in_handle = client.empty(bytes);
29    let out_handle = client.empty(bytes);
30
31    let sample = Box::new(move |iterations: usize| {
32        let start = cubecl_common::profile::Instant::now();
33        unsafe {
34            memory_direct_throughput::launch_unchecked(
35                &client,
36                CubeCount::Static(config.cube_count as u32, 1, 1),
37                CubeDim::new(&client, config.cube_dim),
38                config.vector_size,
39                BufferArg::from_raw_parts(in_handle.clone(), num_lines),
40                BufferArg::from_raw_parts(out_handle.clone(), num_lines),
41                iterations,
42                dtype.into(),
43            )
44        };
45        let _ = cubecl_core::future::block_on(client.sync());
46        start.elapsed()
47    });
48
49    let ops_count = 2 * num_lines * config.vector_size;
50
51    KernelConfig { sample, ops_count }
52}
53
54#[cube(launch_unchecked)]
55pub fn memory_direct_throughput<I: Numeric, N: Size>(
56    input: &[Vector<I, N>],
57    output: &mut [Vector<I, N>],
58    n_iter: usize,
59    #[define(I)] _dtype: StorageType,
60) {
61    let len = output.len();
62    let stride = CUBE_DIM as usize * CUBE_COUNT;
63
64    let steps = (len - ABSOLUTE_POS).div_ceil(stride).max(1);
65
66    for _ in 0..n_iter {
67        for step in 0..steps {
68            let idx = ABSOLUTE_POS + (step * stride);
69
70            if idx < len {
71                output[idx] = input[idx];
72            }
73        }
74    }
75}