use super::SourceTemplate;
use crate::tensor::CubeTensor;
use cubecl::{
CubeKernel, PrecompiledSource,
ir::{UIntKind, metadata::Info},
prelude::*,
};
pub trait KernelSource: Send + 'static + Sync {
fn source(&self) -> SourceTemplate;
fn id(&self) -> KernelId;
fn lang(&self) -> &'static str;
}
#[derive(new)]
pub struct SourceKernel<K> {
kernel_source: K,
cube_dim: CubeDim,
}
impl<K: KernelSource> CubeKernel for SourceKernel<K> {
fn define(&self) -> KernelDefinition {
let settings =
KernelSettings::new(self.cube_dim.0, ExecutionMode::Checked, AddressType::U32);
KernelDefinition {
body: Scope::root(settings.clone()),
settings,
info: Info::default(),
}
}
fn source(&self) -> Option<PrecompiledSource> {
Some(PrecompiledSource {
source: self.kernel_source.source().complete(),
entrypoint_name: "main".to_string(),
lang: self.kernel_source.lang(),
})
}
}
impl<K: KernelSource> KernelMetadata for SourceKernel<K> {
fn id(&self) -> KernelId {
self.kernel_source.id()
}
fn address_type(&self) -> ElemType {
UIntKind::U32.into()
}
}
#[macro_export]
macro_rules! kernel_source {
(
$struct:ident,
$file:expr
) => {
#[derive(new)]
pub struct $struct;
impl $struct {
fn source(&self) -> $crate::template::SourceTemplate {
$crate::template::SourceTemplate::new(include_str!($file))
}
}
};
}
pub fn build_info(tensors: &[&CubeTensor]) -> Vec<u32> {
let ndims = tensors[0].meta.num_dims();
let mut info: Vec<u32> = vec![0; tensors.len() * 2 * ndims + 1];
info[0] = ndims as u32;
let mut current = 1;
for tensor in tensors.iter() {
for d in 0..ndims {
info[current] = tensor.meta.strides()[d] as u32;
current += 1;
}
}
for tensor in tensors.iter() {
for d in 0..ndims {
info[current] = tensor.meta.shape()[d] as u32;
current += 1;
}
}
info
}