Skip to main content

cubecl_std/tensor/contiguous/
base.rs

1use crate::{
2    FastDivmod,
3    tensor::{
4        TensorHandle, into_contiguous,
5        layout::{
6            Layout, LayoutExpand,
7            linear::{LinearLayout, LinearView, linear_layout, linear_view},
8        },
9    },
10};
11use cubecl::prelude::*;
12use cubecl_core::{
13    self as cubecl, calculate_cube_count_elemwise,
14    ir::{StorageType, VectorSize},
15    tensor_vector_size_parallel,
16    zspace::{Strides, strides},
17};
18
19pub const NUM_SM_APPROX: u32 = 50;
20
21/// Returns the offset of the tensor corresponding to the layout tensor.
22#[cube]
23pub fn index_offset_with_layout<T: Scalar, N1: Size, L: Scalar, N2: Size>(
24    tensor: &Tensor<Vector<T, N1>>,
25    layout: &Tensor<Vector<L, N2>>,
26    offset_layout: usize,
27    dim_start: usize,
28    dim_end: usize,
29    #[comptime] unroll: bool,
30) -> usize {
31    let offset_ref = offset_layout * tensor.vector_size();
32    let mut offset = 0;
33
34    #[unroll(unroll)]
35    for i in dim_start..dim_end {
36        let ogwl = offset_ref / layout.stride(i);
37        offset += ogwl % tensor.shape(i) * tensor.stride(i);
38    }
39
40    offset / tensor.vector_size()
41}
42
43/// Returns the offset of the tensor corresponding to a contiguous layout.
44#[cube]
45pub fn index_offset_contiguous<T: Scalar, N: Size>(
46    tensor: &Tensor<Vector<T, N>>,
47    offset_layout: usize,
48    #[comptime] rank: Option<usize>,
49) -> usize {
50    let unroll = rank.is_some();
51    let rank = rank.unwrap_or_else(|| tensor.rank());
52
53    let offset_ref = offset_layout * tensor.vector_size();
54    let mut offset = 0;
55    let mut remainder = offset_ref;
56
57    #[unroll(unroll)]
58    for i in 0..rank {
59        let dim = rank - i - 1;
60        let shape = tensor.shape(dim);
61        let ogwl = remainder % shape;
62        offset += ogwl * tensor.stride(dim);
63        remainder /= shape;
64    }
65
66    offset / tensor.vector_size()
67}
68
69/// Returns the offset of the tensor corresponding to a contiguous layout.
70#[cube]
71pub fn index_offset_contiguous_fastdivmod(
72    offset: usize,
73    shape: &Sequence<FastDivmod<usize>>,
74    stride: &Sequence<usize>,
75    #[comptime] vector_size: VectorSize,
76) -> usize {
77    let rank = shape.len().comptime();
78
79    let offset_ref = offset * vector_size;
80    let mut offset = 0;
81    let mut remainder = offset_ref;
82
83    #[unroll]
84    for i in 0..rank {
85        let dim = rank - i - 1;
86
87        let (rem, ogwl) = shape[dim].div_mod(remainder);
88        offset += ogwl * stride[dim];
89        remainder = rem;
90    }
91
92    offset / vector_size
93}
94
95#[cube(launch, address_type = "dynamic")]
96fn copy_kernel<T: Numeric, N: Size>(
97    input: LinearView<'_, Vector<T, N>>,
98    output: &mut [Vector<T, N>],
99    out_layout: LinearLayout,
100    #[comptime] elems_per_thread: usize,
101    #[define(T)] _elem: StorageType,
102) {
103    let offset_linear = ABSOLUTE_POS * elems_per_thread;
104
105    let mut registers = Array::new(elems_per_thread);
106
107    #[unroll]
108    for i in 0..elems_per_thread {
109        registers[i] = input.read_checked(offset_linear + i);
110    }
111
112    let offset_output = out_layout.to_source_pos(offset_linear);
113
114    #[unroll]
115    for i in 0..elems_per_thread {
116        write_checked(output, offset_output + i, registers[i]);
117    }
118}
119
120#[cube(launch, address_type = "dynamic")]
121fn copy_kernel_pack<T: Numeric, N: Size>(
122    input: LinearView<'_, T>,
123    output: &mut [Vector<T, N>],
124    out_layout: LinearLayout,
125    #[comptime] elems_per_thread: usize,
126    #[define(T)] _elem: StorageType,
127) {
128    let vector_size = output.vector_size().comptime();
129    let vectors_per_thread = elems_per_thread / vector_size;
130
131    let offset_output = ABSOLUTE_POS * vectors_per_thread;
132    let offset_input = offset_output * vector_size;
133
134    let mut registers = Array::new(vectors_per_thread);
135
136    #[unroll]
137    for i in 0..vectors_per_thread {
138        let offset = i * vector_size;
139        let mut reg = Vector::<T, N>::empty();
140        #[unroll]
141        for k in 0..vector_size {
142            let offset_input = offset_input + offset + k;
143            reg.insert(k, input.read_checked(offset_input));
144        }
145        registers[i] = reg;
146    }
147
148    let offset_output = out_layout.to_source_pos(offset_output);
149
150    #[unroll]
151    for i in 0..vectors_per_thread {
152        write_checked(output, offset_output + i, registers[i]);
153    }
154}
155
156/// Fetch all values required contained in a given position, unpack them, then repack them to their
157/// new position.
158#[cube]
159fn index_packed<N: Int>(
160    tensor: &Tensor<N>,
161    pos: usize,
162    in_shape: &Sequence<FastDivmod<usize>>,
163    #[comptime] packed_dim: usize,
164    #[comptime] packing: usize,
165    #[comptime] rank: usize,
166) -> N {
167    let type_size_bits = N::type_size_bits().comptime();
168    let bits_per_elem = type_size_bits / packing;
169    let mask = (1u32 << bits_per_elem) - 1;
170    let mask = N::cast_from(mask);
171
172    let elem_pos = pos * packing;
173
174    let mut out = N::new(0);
175    for n in 0..packing {
176        let mut remainder = elem_pos + n;
177        let mut offset = 0;
178        let mut packing_offset = 0;
179
180        #[unroll]
181        for i in 0..rank {
182            let dim = rank - i - 1;
183            let (rem, mut local_pos) = in_shape[dim].div_mod(remainder);
184            remainder = rem;
185            if dim == packed_dim {
186                packing_offset = local_pos % packing;
187                local_pos /= packing;
188            }
189            offset += local_pos * tensor.stride(dim);
190        }
191        let packed_val = tensor[offset];
192        let shift_in = packing_offset * bits_per_elem;
193        let shift_out = n * bits_per_elem;
194        let value = (packed_val >> N::cast_from(shift_in)) & mask;
195
196        out |= value << N::cast_from(shift_out);
197    }
198    out
199}
200
201#[cube(launch, address_type = "dynamic")]
202fn copy_kernel_packed<T: Int, N: Size>(
203    input: &Tensor<T>,
204    output: &mut Tensor<Vector<T, N>>,
205    out_layout: LinearLayout,
206    in_shape: Sequence<FastDivmod<usize>>,
207    #[comptime] packed_dim: usize,
208    #[comptime] packing: usize,
209    #[comptime] rank: usize,
210    #[comptime] elems_per_thread: usize,
211    #[define(T)] _elem: StorageType,
212) {
213    let vector_size = output.vector_size().comptime();
214    let vectors_per_thread = elems_per_thread / vector_size;
215
216    let offset_output = ABSOLUTE_POS * vectors_per_thread;
217    let offset_input = offset_output * vector_size;
218
219    if offset_output >= output.len() {
220        terminate!()
221    }
222
223    let mut registers = Array::new(vectors_per_thread);
224
225    #[unroll]
226    for i in 0..vectors_per_thread {
227        let offset = i * vector_size;
228        let mut reg = Vector::<T, N>::empty();
229        #[unroll]
230        for k in 0..vector_size {
231            let offset_input = offset_input + offset + k;
232
233            reg.insert(
234                k,
235                index_packed(input, offset_input, &in_shape, packed_dim, packing, rank),
236            );
237        }
238        registers[i] = reg;
239    }
240
241    let offset_output = out_layout.to_source_pos(offset_output);
242
243    #[unroll]
244    for i in 0..vectors_per_thread {
245        output[offset_output + i] = registers[i];
246    }
247}
248
249/// Make a jit tensor contiguous, using the pitched allocator if available.
250/// See [`create_tensor`](cubecl_runtime::client::ComputeClient::create_tensor).
251/// Handles unpacking and repacking packed tensors (i.e. quantized values).
252/// `shape` refers to the actual (unpacked) shape of the tensor, while `packing` specifies the
253/// number of elements in each storage element.
254///
255/// # Warning
256/// This assumes `u32` or `u8` packing.
257pub fn into_contiguous_packed<R: Runtime>(
258    client: &ComputeClient<R>,
259    input: TensorBinding<R>,
260    packed_dim: usize,
261    shape: &[usize],
262    packing: usize,
263    dtype: StorageType,
264) -> TensorHandle<R> {
265    let rank = shape.len();
266    if rank <= 1 {
267        return into_contiguous(client, input, dtype);
268    }
269
270    let mut out_shape = shape.to_vec();
271    out_shape[rank - 1] = out_shape[rank - 1].div_ceil(packing);
272    let output = TensorHandle::empty(client, out_shape, dtype);
273
274    // Should reinterpret as u8 if possible at some point, but requires modifying shape/strides so
275    // keep it simple for now
276    into_contiguous_packed_ref(
277        client,
278        input,
279        output.clone().binding(),
280        packed_dim,
281        shape,
282        packing,
283        dtype,
284    );
285
286    output
287}
288
289/// Make a jit tensor contiguous.
290pub fn copy_gpu_ref<R: Runtime>(
291    client: &ComputeClient<R>,
292    input: TensorBinding<R>,
293    output: TensorBinding<R>,
294    dtype: StorageType,
295) {
296    let num_elems: usize = input.shape.iter().product();
297
298    // Vectorization is only enabled when the last dimension is contiguous.
299    let in_rank = input.strides.len();
300    let out_rank = output.strides.len();
301    let vector_size_in = tensor_vector_size_parallel(
302        client.io_optimized_vector_sizes(dtype.size()),
303        &input.shape,
304        &input.strides,
305        in_rank - 1,
306    );
307    let vector_size_out = tensor_vector_size_parallel(
308        client.io_optimized_vector_sizes(dtype.size()),
309        &output.shape,
310        &output.strides,
311        out_rank - 1,
312    );
313    let vector_size = vector_size_in.min(vector_size_out);
314
315    let num_vecs = num_elems / vector_size as usize;
316    let num_sm = client
317        .properties()
318        .hardware
319        .num_streaming_multiprocessors
320        .unwrap_or(NUM_SM_APPROX);
321    let cube_dim = CubeDim::new(client, num_vecs);
322    let simul_vecs = num_sm * cube_dim.num_elems();
323    let mut elems_per_unit = match num_vecs / simul_vecs as usize {
324        0..2 => 1,
325        2..4 => 2,
326        4..8 => 4,
327        8.. => 8,
328    };
329
330    let mut num_elems_per_unit = vector_size as usize * elems_per_unit;
331
332    let last_dim = output.shape[out_rank - 1];
333
334    // If tensor is strided, elems_per_unit must be compatible with last dim
335    while !last_dim.is_multiple_of(num_elems_per_unit as usize) {
336        elems_per_unit /= 2;
337        num_elems_per_unit /= 2;
338    }
339
340    let out_vec = if vector_size > 1 {
341        vector_size
342    } else {
343        // Recompute because it needs to account for `num_elems_per_unit`
344        client
345            .io_optimized_vector_sizes(dtype.size())
346            .filter(|it| num_elems_per_unit.is_multiple_of(*it))
347            .max()
348            .unwrap_or(1)
349    };
350
351    let address_type = input
352        .required_address_type(dtype.size())
353        .max(output.required_address_type(dtype.size()));
354    let input = linear_view(input);
355    let out_layout = linear_layout(&output, out_vec);
356
357    let cube_count = calculate_cube_count_elemwise(
358        client,
359        num_elems.div_ceil(num_elems_per_unit as usize),
360        cube_dim,
361    );
362
363    let launch = if vector_size != out_vec && out_vec > 1 {
364        copy_kernel_pack::launch
365    } else {
366        copy_kernel::launch
367    };
368
369    launch(
370        client,
371        cube_count,
372        cube_dim,
373        address_type,
374        out_vec,
375        input,
376        output.clone().into_buffer_arg(),
377        out_layout,
378        elems_per_unit,
379        dtype,
380    )
381}
382
383/// Make a jit tensor contiguous.
384pub fn into_contiguous_packed_ref<R: Runtime>(
385    client: &ComputeClient<R>,
386    input: TensorBinding<R>,
387    output: TensorBinding<R>,
388    packed_dim: usize,
389    shape: &[usize],
390    packing: usize,
391    dtype: StorageType,
392) {
393    let num_elems: usize = input.shape.iter().product();
394
395    // Vectorization is only enabled when the last dimension is contiguous.
396    let in_rank = input.strides.len();
397    let out_rank = output.strides.len();
398    let in_packed_dim = in_rank - packed_dim - 1;
399    let vector_size = tensor_vector_size_parallel(
400        client.io_optimized_vector_sizes(dtype.size()),
401        &output.shape,
402        &output.strides,
403        out_rank - 1,
404    );
405    let num_vecs = num_elems / vector_size as usize;
406    let num_sm = client
407        .properties()
408        .hardware
409        .num_streaming_multiprocessors
410        .unwrap_or(NUM_SM_APPROX);
411
412    let cube_dim = CubeDim::new(client, num_vecs);
413    let simul_vecs = num_sm * cube_dim.num_elems();
414    let elems_per_unit = match num_vecs / simul_vecs as usize {
415        0..2 => 1,
416        2..4 => 2,
417        4..8 => 4,
418        8.. => 8,
419    };
420
421    let mut num_elems_per_unit = vector_size as usize * elems_per_unit;
422
423    let last_dim = output.shape[out_rank - 1];
424
425    // If tensor is strided, num_elems_per_unit must be compatible with last dim
426    while !last_dim.is_multiple_of(num_elems_per_unit as usize) {
427        num_elems_per_unit /= 2;
428    }
429
430    let out_layout = linear_layout(&output, vector_size);
431
432    let address_type = input
433        .required_address_type(dtype.size())
434        .max(output.required_address_type(dtype.size()));
435    let cube_count = calculate_cube_count_elemwise(
436        client,
437        num_elems.div_ceil(num_elems_per_unit as usize),
438        cube_dim,
439    );
440
441    let in_shape = shape.iter().copied().collect();
442
443    copy_kernel_packed::launch(
444        client,
445        cube_count,
446        cube_dim,
447        address_type,
448        vector_size,
449        input.into_tensor_arg(),
450        output.into_tensor_arg(),
451        out_layout,
452        in_shape,
453        in_packed_dim,
454        packing,
455        in_rank,
456        num_elems_per_unit,
457        dtype,
458    )
459}
460
461/// Checks if the tensor associated with the given shape and strides is contiguous.
462pub fn is_contiguous(shape: &[usize], strides: &[usize]) -> bool {
463    if shape.is_empty() {
464        return true;
465    }
466
467    for (&expected, &stride) in compact_strides(shape).iter().zip(strides) {
468        if expected != stride {
469            return false;
470        }
471    }
472
473    true
474}
475
476/// Checks if a tensor is only strided on the last dimension, and could be safely reinterpreted as
477/// a 2D tensor with unit stride on the last dimension. This will always hold for non-permuted
478/// tensors allocated on a runtime.
479pub fn is_contiguous_pitched(shape: &[usize], strides: &[usize]) -> bool {
480    let rank = shape.len();
481    if strides[rank - 1] != 1 {
482        return false;
483    }
484    if rank <= 1 {
485        return true;
486    }
487
488    let mut sorted = strides.to_vec();
489    sorted.sort();
490    sorted.reverse();
491
492    if sorted != strides {
493        return false;
494    }
495
496    for i in 0..rank - 2 {
497        if strides[i] != shape[i + 1] * strides[i + 1] {
498            return false;
499        }
500    }
501    true
502}
503
504pub fn compact_strides(shape: &[usize]) -> Strides {
505    let rank = shape.len();
506    let mut strides = strides![1; rank];
507    for i in (0..rank - 1).rev() {
508        strides[i] = strides[i + 1] * shape[i + 1];
509    }
510    strides
511}