use cubecl::prelude::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct StorageAlign(u64);
impl StorageAlign {
pub(crate) fn from_client<R: Runtime>(client: &ComputeClient<R>) -> Self {
Self::new(client.properties().memory.alignment)
}
pub(crate) fn new(bytes: u64) -> Self {
debug_assert!(
bytes.is_power_of_two(),
"storage alignment {bytes} is not a power of two"
);
Self(bytes.max(1))
}
pub(crate) fn pad_bytes(self, bytes: u64) -> u64 {
bytes.next_multiple_of(self.0)
}
pub(crate) fn pad_elems<T>(self, elems: usize) -> usize {
let per_boundary = (self.0 as usize).div_ceil(size_of::<T>()).max(1);
elems.next_multiple_of(per_boundary)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pad_bytes_rounds_up_to_the_boundary() {
let align = StorageAlign::new(32);
assert_eq!(align.pad_bytes(0), 0);
assert_eq!(align.pad_bytes(1), 32);
assert_eq!(align.pad_bytes(32), 32);
assert_eq!(align.pad_bytes(33), 64);
}
#[test]
fn pad_elems_spans_whole_boundaries() {
let align = StorageAlign::new(32);
assert_eq!(align.pad_elems::<f32>(0), 0);
assert_eq!(align.pad_elems::<f32>(1), 8);
assert_eq!(align.pad_elems::<f32>(8), 8);
assert_eq!(align.pad_elems::<f32>(9), 16);
}
#[test]
fn pad_elems_tracks_a_larger_alignment() {
let align = StorageAlign::new(256);
assert_eq!(align.pad_elems::<f32>(1), 64);
assert_eq!(align.pad_elems::<f32>(64), 64);
assert_eq!(align.pad_elems::<f32>(65), 128);
}
#[test]
fn padded_element_counts_are_byte_aligned() {
for bytes in [4u64, 16, 32, 64, 256] {
let align = StorageAlign::new(bytes);
for elems in [1usize, 3, 7, 137, 24_660] {
let padded = align.pad_elems::<f32>(elems) as u64 * size_of::<f32>() as u64;
assert_eq!(padded % bytes, 0, "align {bytes}, {elems} elements");
}
}
}
#[test]
fn an_alignment_below_the_element_size_leaves_counts_alone() {
let align = StorageAlign::new(4);
assert_eq!(align.pad_elems::<f32>(3), 3);
}
}