cubek-std 0.3.0-pre.1

CubeK: Standard Library
Documentation
use cubecl::prelude::*;

use crate::{
    MatrixLayout, StageIdent, TileSize,
    tile::{SharedTile, Tile, TileKind, TileKindExpand, TileScope},
};

/// Interleaved-on-k tile. Holds just the minimal comptime data the body uses
/// (`tile_size` and `plane_dim`); the matmul-level configuration (and any
/// metadata not consumed by the tile body, like swizzle modes) lives in
/// cubek-matmul as `InterleavedMatmul`.
#[derive(CubeType)]
pub struct InterleavedTile<N: Numeric> {
    pub data: Array<N>,
    #[cube(comptime)]
    pub matrix_layout: MatrixLayout,
    #[cube(comptime)]
    pub tile_size: TileSize,
    #[cube(comptime)]
    pub plane_dim: u32,
}

#[cube]
pub fn interleaved_allocate_lhs<L: Numeric, Sc: TileScope>(
    #[comptime] layout: MatrixLayout,
    #[comptime] tile_size: TileSize,
    #[comptime] plane_dim: u32,
) -> Tile<L, Sc> {
    let m = tile_size.m();
    let k = tile_size.k();
    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<L> {
        data: Array::new((m * (k / plane_dim)) as usize),
        matrix_layout: layout,
        tile_size,
        plane_dim,
    }))
}

#[cube]
pub fn interleaved_allocate_rhs<R: Numeric, Sc: TileScope>(
    #[comptime] layout: MatrixLayout,
    #[comptime] tile_size: TileSize,
    #[comptime] plane_dim: u32,
) -> Tile<R, Sc> {
    let n = tile_size.n();
    let k = tile_size.k();
    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<R> {
        data: Array::new(((k / plane_dim) * n) as usize),
        matrix_layout: layout,
        tile_size,
        plane_dim,
    }))
}

#[cube]
pub fn interleaved_allocate_acc<A: Numeric, Sc: TileScope>(
    #[comptime] layout: MatrixLayout,
    #[comptime] tile_size: TileSize,
    #[comptime] plane_dim: u32,
) -> Tile<A, Sc> {
    let m = tile_size.m();
    let n = tile_size.n();
    Tile::from_kind(TileKind::new_Interleaved(InterleavedTile::<A> {
        data: Array::new((m * n) as usize),
        matrix_layout: layout,
        tile_size,
        plane_dim,
    }))
}

#[cube]
impl<A: Numeric> InterleavedTile<A> {
    /// Executes `lhs ยท rhs`, accumulating into `self` via the plane-
    /// interleaved-on-k matmul.
    pub fn mma<L: Numeric, R: Numeric>(
        &mut self,
        lhs: &InterleavedTile<L>,
        rhs: &InterleavedTile<R>,
    ) {
        interleaved_execute(
            &lhs.data,
            lhs.matrix_layout,
            &rhs.data,
            rhs.matrix_layout,
            &mut self.data,
            self.matrix_layout,
            self.tile_size,
            self.plane_dim,
        );
    }
}

#[cube]
impl<N: Numeric> InterleavedTile<N> {
    /// Copies into the interleaved tile from `source`. Supported sources:
    /// `Shared` and `None` (zero-init).
    pub fn copy_from<SE: Numeric, SS: Size, Sc: TileScope>(
        &mut self,
        source: &Tile<SE, Sc>,
        #[comptime] ident: StageIdent,
    ) {
        match &source.kind {
            TileKind::SharedTile(shared) => {
                interleaved_load_from_shared::<SE, SS, N>(
                    shared,
                    &mut self.data,
                    self.tile_size,
                    self.plane_dim,
                    ident,
                );
            }
            TileKind::None => {
                interleaved_load_zeros::<N>(&mut self.data, self.tile_size);
            }
            TileKind::Cmma(_)
            | TileKind::Mma(_)
            | TileKind::Register(_)
            | TileKind::PlaneVec(_)
            | TileKind::Interleaved(_)
            | TileKind::Unit(_)
            | TileKind::WhiteboxFragment(_)
            | TileKind::RowWise(_)
            | TileKind::Bounce(_)
            | TileKind::Stage(_)
            | TileKind::Partition(_)
            | TileKind::Pipelined(_) => {
                panic!("InterleavedTile::copy_from: unsupported source variant")
            }
        }
    }

    pub fn init_zero(&mut self) {
        interleaved_load_zeros::<N>(&mut self.data, self.tile_size);
    }
}

// ===========================================================================
// Compute: matmul / load / write / zero-init
// ===========================================================================

#[cube]
#[allow(clippy::too_many_arguments)]
pub fn interleaved_execute<L: Numeric, R: Numeric, A: Numeric>(
    lhs: &Array<L>,
    #[comptime] lhs_layout: MatrixLayout,
    rhs: &Array<R>,
    #[comptime] rhs_layout: MatrixLayout,
    acc: &mut Array<A>,
    #[comptime] _acc_layout: MatrixLayout,
    #[comptime] tile_size: TileSize,
    #[comptime] plane_dim: u32,
) {
    let m = tile_size.m() as usize;
    let n = tile_size.n() as usize;
    let k = tile_size.k() as usize;
    let plane_dim = plane_dim as usize;
    let local_k = k / plane_dim;

    let (lhs_row_count, lhs_col_count) = (m, local_k);
    let (rhs_row_count, rhs_col_count) = (local_k, n);

    #[unroll]
    for m_ in 0..m {
        #[unroll]
        for n_ in 0..n {
            #[unroll]
            for k_ in 0..local_k {
                let lhs_elem = A::cast_from(match lhs_layout {
                    MatrixLayout::RowMajor => lhs[m_ * lhs_col_count + k_],
                    MatrixLayout::ColMajor => lhs[k_ * lhs_row_count + m_],
                });
                let rhs_elem = A::cast_from(match rhs_layout {
                    MatrixLayout::RowMajor => rhs[k_ * rhs_col_count + n_],
                    MatrixLayout::ColMajor => rhs[n_ * rhs_row_count + k_],
                });
                acc[m_ * n + n_] += lhs_elem * rhs_elem;
            }
        }
    }
}

#[cube]
pub fn interleaved_load_from_shared<E: Numeric, ES: Size, N: Numeric>(
    shared: &SharedTile<E>,
    arr: &mut Array<N>,
    #[comptime] tile_size: TileSize,
    #[comptime] plane_dim: u32,
    #[comptime] ident: StageIdent,
) {
    let shared = shared.view::<ES>();
    let shared = &shared;
    match ident {
        StageIdent::Lhs | StageIdent::Rhs => {
            let m = tile_size.m() as usize;
            let n = tile_size.n() as usize;
            let k = tile_size.k() as usize;
            let plane_dim = plane_dim as usize;
            let k_local = k / plane_dim;

            let shared_layout = comptime!(shared.layout);
            let vector_size = ES::value();

            let unit_id = UNIT_POS_X as usize;
            let k_offset = k_local * unit_id;

            let (strided_dim_count, contiguous_dim_count) = match (shared_layout, ident) {
                (MatrixLayout::RowMajor, StageIdent::Lhs) => (m, k_local),
                (MatrixLayout::RowMajor, StageIdent::Rhs) => (k_local, n),
                (MatrixLayout::ColMajor, StageIdent::Lhs) => (k_local, m),
                (MatrixLayout::ColMajor, StageIdent::Rhs) => (n, k_local),
                _ => unreachable!(),
            };

            let (strided_dim_offset, contiguous_dim_offset) = match (shared_layout, ident) {
                (MatrixLayout::RowMajor, StageIdent::Lhs)
                | (MatrixLayout::ColMajor, StageIdent::Rhs) => (0, k_offset / vector_size),
                (MatrixLayout::RowMajor, StageIdent::Rhs)
                | (MatrixLayout::ColMajor, StageIdent::Lhs) => (k_offset, 0),
                _ => unreachable!(),
            };

            assert!(contiguous_dim_count % vector_size == 0);
            let vector_count_in_dim = contiguous_dim_count / vector_size;

            for i in 0..strided_dim_count {
                for j in 0..vector_count_in_dim {
                    let vector = Vector::<N, ES>::cast_from(shared.get_vector(
                        (i + strided_dim_offset) as u32,
                        (j + contiguous_dim_offset) as u32,
                    ));
                    let vector_start = i * contiguous_dim_count + j * vector_size;
                    for l in 0..vector_size {
                        arr[vector_start + l] = vector.extract(l);
                    }
                }
            }
        }
        StageIdent::Acc => {
            panic!("Not yet implemented: Interleaved acc load from shared");
        }
        _ => panic!("Invalid ident for Interleaved load"),
    }
}

#[cube]
pub fn interleaved_load_zeros<N: Numeric>(arr: &mut Array<N>, #[comptime] tile_size: TileSize) {
    let m = tile_size.m() as usize;
    let n = tile_size.n() as usize;
    let size = m * n;
    for i in 0..size {
        arr[i] = N::from_int(0);
    }
}

#[cube]
pub fn interleaved_write_to_shared<E: Numeric, ES: Size, A: Numeric>(
    shared: &mut SharedTile<E>,
    arr: &Array<A>,
    #[comptime] tile_size: TileSize,
) {
    let mut shared = shared.view::<ES>();
    let shared = &mut shared;
    let m = tile_size.m();
    let n = tile_size.n();
    let out_vector_size = shared.container.vector_size().comptime() as u32;
    let size_mn = m * n;

    // `plane_sum` reduces across the plane, so every unit must participate. Only unit 0 stores.
    #[unroll]
    for i in 0..size_mn / out_vector_size {
        let mut vector = Vector::<A, ES>::empty();
        #[unroll]
        for j in 0..out_vector_size {
            vector.insert(
                j as usize,
                plane_sum(arr[(i * out_vector_size + j) as usize]),
            );
        }
        if UNIT_POS_X == 0 {
            let offs = shared.stage_offset(i);
            shared.container[offs as usize] = Vector::cast_from(vector);
        }
    }
}