use std::sync::Arc;
use teeny_core::{
graph::{DtypeRepr, Graph, Op, Shape},
model::{ExecutableOp, Lowering, LoweringMode, RuntimeOp},
utils::dag::Dag,
};
use crate::nn::{
activation::extra::{
LogSoftmaxBackward, LogSoftmaxForward, PreluForward, ShrinkRuntimeOp, SwishBackward,
SwishForward, ThresholdedReluRuntimeOp,
},
activation::{
elu::{
CeluForward, CeluForwardDispatch, EluForward, EluForwardDispatch, SeluForward,
SeluForwardDispatch,
},
gelu::{GeluForwardDispatch, MishForward, MishForwardDispatch},
hard::{
HardshrinkForward, HardshrinkForwardDispatch, HardsigmoidForward,
HardsigmoidForwardDispatch, HardswishForward, HardswishForwardDispatch,
HardtanhForward, HardtanhForwardDispatch, Relu6Forward, Relu6ForwardDispatch,
},
misc::{
LeakyReluForward, LeakyReluForwardDispatch, SoftplusForward, SoftplusForwardDispatch,
SoftshrinkForward, SoftshrinkForwardDispatch, SoftsignForward, SoftsignForwardDispatch,
ThresholdForward, ThresholdForwardDispatch,
},
relu::{ReluBackward, ReluForward},
sigmoid::{
LogsigmoidForward, LogsigmoidForwardDispatch, SigmoidForwardDispatch,
SiluForwardDispatch,
},
softmax::SoftmaxForward,
tanh::{TanhForward, TanhForwardDispatch, TanhshrinkForward, TanhshrinkForwardDispatch},
},
conv::{
conv1d::Conv1dForward,
conv2d::{Conv2dBackward, Conv2dBiasForward, Conv2dForward},
conv3d::Conv3dForward,
},
fused::{
conv2d_bn_silu::Conv2dBnSiluForward, conv2d_bn_silu_gemm::Conv2dBnSiluGemmForward,
conv2d_bn_silu_tiled::Conv2dBnSiluTiledForward,
},
mlp::{
flatten::FlattenForward,
linear::{LinearBackward, LinearForward},
},
norm::{
batchnorm::{BatchNorm2dNchwInferenceRuntimeOp, BatchNormForwardInference},
groupnorm::GroupNormForwardInference,
instancenorm::InstanceNormForwardInference,
layernorm::{LayerNormForwardInference, LayerNormForwardInferenceRuntimeOp},
rmsnorm::RmsNormForward,
},
pad::{
circular_pad1d::CircularPad1dForward, circular_pad2d::CircularPad2dForward,
circular_pad3d::CircularPad3dForward, constant_pad1d::ConstantPad1dForward,
constant_pad2d::ConstantPad2dForward, constant_pad3d::ConstantPad3dForward,
reflection_pad1d::ReflectionPad1dForward, reflection_pad2d::ReflectionPad2dForward,
reflection_pad3d::ReflectionPad3dForward, replication_pad1d::ReplicationPad1dForward,
replication_pad2d::ReplicationPad2dForward, replication_pad3d::ReplicationPad3dForward,
},
pool::{
avgpool1d::Avgpool1dForward,
avgpool2d::Avgpool2dForward,
avgpool3d::Avgpool3dForward,
lppool1d::Lppool1dForward,
lppool2d::Lppool2dForward,
lppool3d::Lppool3dForward,
maxpool1d::Maxpool1dForward,
maxpool2d::{Maxpool2dBackward, Maxpool2dForward},
maxpool3d::Maxpool3dForward,
},
tensor::{
channel_bias_add::{ChannelBiasAddRuntimeOp, NchwBiasAddRuntimeOp},
channel_cat::ChannelCatRuntimeOp,
channel_chunk::ChannelChunkRuntimeOp,
elemwise_add::{ElemwiseAddBackward, ElemwiseAddForward},
elemwise_binary::{
ClipRuntimeOp, ElemwiseDivBackward, ElemwiseDivForward, ElemwiseEqualForward,
ElemwiseFmodForward, ElemwiseGreaterEqualForward, ElemwiseGreaterForward,
ElemwiseLessEqualForward, ElemwiseLessForward, ElemwiseMaxBackward, ElemwiseMaxForward,
ElemwiseMeanBackward, ElemwiseMeanForward, ElemwiseMinBackward, ElemwiseMinForward,
ElemwiseMulBackward, ElemwiseMulForward, ElemwisePowBackward, ElemwisePowForward,
ElemwiseSubBackward, ElemwiseSubForward, ElemwiseSumBackward, ElemwiseSumForward,
ElemwiseWhereBackward, ElemwiseWhereForward,
},
elemwise_unary::{
ElemwiseAbsBackward, ElemwiseAbsForward, ElemwiseAcosBackward, ElemwiseAcosForward,
ElemwiseAcoshBackward, ElemwiseAcoshForward, ElemwiseAsinBackward, ElemwiseAsinForward,
ElemwiseAsinhBackward, ElemwiseAsinhForward, ElemwiseAtanBackward, ElemwiseAtanForward,
ElemwiseAtanhBackward, ElemwiseAtanhForward, ElemwiseCeilForward, ElemwiseCosBackward,
ElemwiseCosForward, ElemwiseCoshBackward, ElemwiseCoshForward, ElemwiseErfBackward,
ElemwiseErfForward, ElemwiseExpBackward, ElemwiseExpForward, ElemwiseFloorForward,
ElemwiseIsnanForward, ElemwiseLogBackward, ElemwiseLogForward, ElemwiseNegBackward,
ElemwiseNegForward, ElemwiseReciprocalBackward, ElemwiseReciprocalForward,
ElemwiseSignForward, ElemwiseSinBackward, ElemwiseSinForward, ElemwiseSinhBackward,
ElemwiseSinhForward, ElemwiseSqrtBackward, ElemwiseSqrtForward, ElemwiseTanBackward,
ElemwiseTanForward,
},
reduction::{
CumProdForward, CumSumForward, GlobalAvgPoolForward, GlobalMaxPoolForward,
ReduceL1Forward, ReduceL2Forward, ReduceLogSumExpForward, ReduceLogSumForward,
ReduceMaxForward, ReduceMeanForward, ReduceMinForward, ReduceProdForward,
ReduceSumForward, ReduceSumSquareForward,
},
upsample_nearest2d::{UpsampleNearest2dBackward, UpsampleNearest2dForward},
},
};
use crate::math::gemm::MatMulRuntimeOp;
use crate::errors::Result;
#[cfg(feature = "training")]
use crate::nn::norm::batchnorm::{
BatchNorm2dNchwBackward, BatchNormNormalizeForward, BatchNormNormalizeRuntimeOp,
BatchNormStatsForward, BatchNormStatsRuntimeOp,
};
macro_rules! make_num_kernel {
($K:ident ($($arg:expr),*), $node:expr) => {{
let (name, ks, rop) = match $node.dtype {
DtypeRepr::F32 => { let k = $K::<f32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::F64 => { let k = $K::<f64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I8 => { let k = $K::<i8>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I16 => { let k = $K::<i16>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I32 => { let k = $K::<i32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I64 => { let k = $K::<i64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U8 => { let k = $K::<u8>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U16 => { let k = $K::<u16>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U32 => { let k = $K::<u32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U64 => { let k = $K::<u64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
other => return Err(anyhow::anyhow!("{:?} is not a supported Num dtype for {}", other, stringify!($K))),
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: $node.shape.clone(),
dtype: $node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}};
($K:ident ($($arg:expr),*), $Bwd:ident ($($barg:expr),*), $node:expr) => {{
let (name, ks, rop) = match $node.dtype {
DtypeRepr::F32 => { let k = $K::<f32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::F64 => { let k = $K::<f64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I8 => { let k = $K::<i8>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I16 => { let k = $K::<i16>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I32 => { let k = $K::<i32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::I64 => { let k = $K::<i64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U8 => { let k = $K::<u8>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U16 => { let k = $K::<u16>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U32 => { let k = $K::<u32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::U64 => { let k = $K::<u64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
other => return Err(anyhow::anyhow!("{:?} is not a supported Num dtype for {}", other, stringify!($K))),
};
#[cfg(feature = "training")]
let (bwd_name, bwd_ks) = match $node.dtype {
DtypeRepr::F32 => { let k = $Bwd::<f32>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::F64 => { let k = $Bwd::<f64>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::I8 => { let k = $Bwd::<i8>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::I16 => { let k = $Bwd::<i16>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::I32 => { let k = $Bwd::<i32>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::I64 => { let k = $Bwd::<i64>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::U8 => { let k = $Bwd::<u8>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::U16 => { let k = $Bwd::<u16>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::U32 => { let k = $Bwd::<u32>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::U64 => { let k = $Bwd::<u64>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
other => return Err(anyhow::anyhow!("{:?} is not a supported Num dtype for {}", other, stringify!($Bwd))),
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: $node.shape.clone(),
dtype: $node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: format!("{}_entry_point", bwd_name),
runtime_op: rop,
})
}};
}
macro_rules! make_float_kernel {
($K:ident ($($arg:expr),*), $node:expr) => {{
let (name, ks, rop) = match $node.dtype {
DtypeRepr::F32 => { let k = $K::<f32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::F64 => { let k = $K::<f64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
other => return Err(anyhow::anyhow!("{:?} is not a Float dtype for {}", other, stringify!($K))),
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: $node.shape.clone(),
dtype: $node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}};
($K:ident ($($arg:expr),*), $Bwd:ident ($($barg:expr),*), $node:expr) => {{
let (name, ks, rop) = match $node.dtype {
DtypeRepr::F32 => { let k = $K::<f32>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
DtypeRepr::F64 => { let k = $K::<f64>::new($($arg),*); let nm = k.name.to_string(); let src = k.source.clone(); let r: Arc<dyn RuntimeOp> = Arc::new(k); (nm, src, r) }
other => return Err(anyhow::anyhow!("{:?} is not a Float dtype for {}", other, stringify!($K))),
};
#[cfg(feature = "training")]
let (bwd_name, bwd_ks) = match $node.dtype {
DtypeRepr::F32 => { let k = $Bwd::<f32>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
DtypeRepr::F64 => { let k = $Bwd::<f64>::new($($barg),*); (k.name.to_string(), k.source.clone()) }
other => return Err(anyhow::anyhow!("{:?} is not a Float dtype for {}", other, stringify!($Bwd))),
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: $node.shape.clone(),
dtype: $node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: format!("{}_entry_point", bwd_name),
runtime_op: rop,
})
}};
}
fn exec_from(
shape: Shape,
dtype: DtypeRepr,
inst: teeny_core::model::KernelInstance,
) -> Box<KernelExecutable> {
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", inst.name),
name: inst.name,
kernel_source: inst.source,
shape,
dtype,
#[cfg(feature = "training")]
backward_kernel_source: inst
.backward
.as_ref()
.map(|b| b.source.clone())
.unwrap_or_default(),
#[cfg(feature = "training")]
backward_entry_point: inst
.backward
.as_ref()
.map(|b| format!("{}_entry_point", b.name))
.unwrap_or_default(),
runtime_op: inst.runtime_op,
})
}
pub struct KernelExecutable {
pub name: String,
pub kernel_source: String,
pub entry_point: String,
pub shape: Shape,
pub dtype: DtypeRepr,
pub runtime_op: Arc<dyn RuntimeOp>,
#[cfg(feature = "training")]
pub backward_kernel_source: String,
#[cfg(feature = "training")]
pub backward_entry_point: String,
}
impl ExecutableOp for KernelExecutable {
fn name(&self) -> &str {
&self.name
}
fn is_input(&self) -> bool {
self.name == "input"
}
fn forward_kernel_source(&self) -> &str {
&self.kernel_source
}
fn forward_kernel_entry_point(&self) -> &str {
&self.entry_point
}
fn output_shape(&self) -> &Shape {
&self.shape
}
fn output_dtype(&self) -> DtypeRepr {
self.dtype
}
fn runtime_op(&self) -> Option<Arc<dyn RuntimeOp>> {
if self.is_input() {
None
} else {
Some(Arc::clone(&self.runtime_op))
}
}
#[cfg(feature = "training")]
fn backward_kernel_source(&self) -> &str {
&self.backward_kernel_source
}
#[cfg(feature = "training")]
fn backward_kernel_entry_point(&self) -> &str {
&self.backward_entry_point
}
}
macro_rules! impl_stub_runtime_op_num {
($T:ident) => {
impl<D: teeny_core::dtype::Num + Send + Sync + 'static> RuntimeOp for $T<D> {
fn n_activation_inputs(&self) -> usize {
unimplemented!(concat!(stringify!($T), " has no runtime support"))
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
unimplemented!()
}
fn pack_args(
&self,
_: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
_: teeny_core::model::RawPtr,
_: &[usize],
_: i32,
_: &mut dyn teeny_core::device::program::ArgVisitor,
) {
unimplemented!()
}
fn block(&self) -> [u32; 3] {
unimplemented!()
}
fn grid(&self, _: &[usize]) -> [u32; 3] {
unimplemented!()
}
}
};
}
macro_rules! impl_stub_runtime_op_float {
($T:ident) => {
impl<D: teeny_core::dtype::Float + Send + Sync + 'static> RuntimeOp for $T<D> {
fn n_activation_inputs(&self) -> usize {
unimplemented!(concat!(stringify!($T), " has no runtime support"))
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
unimplemented!()
}
fn pack_args(
&self,
_: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
_: teeny_core::model::RawPtr,
_: &[usize],
_: i32,
_: &mut dyn teeny_core::device::program::ArgVisitor,
) {
unimplemented!()
}
fn block(&self) -> [u32; 3] {
unimplemented!()
}
fn grid(&self, _: &[usize]) -> [u32; 3] {
unimplemented!()
}
}
};
}
impl_stub_runtime_op_float!(BatchNormForwardInference);
impl_stub_runtime_op_float!(LayerNormForwardInference);
impl_stub_runtime_op_float!(RmsNormForward);
impl_stub_runtime_op_float!(GroupNormForwardInference);
impl_stub_runtime_op_float!(InstanceNormForwardInference);
impl_stub_runtime_op_num!(Conv3dForward);
impl_stub_runtime_op_num!(Avgpool1dForward);
impl_stub_runtime_op_num!(Avgpool3dForward);
impl_stub_runtime_op_num!(Maxpool1dForward);
impl_stub_runtime_op_num!(Maxpool3dForward);
impl_stub_runtime_op_float!(Lppool1dForward);
impl_stub_runtime_op_float!(Lppool2dForward);
impl_stub_runtime_op_float!(Lppool3dForward);
impl_stub_runtime_op_num!(ConstantPad1dForward);
impl_stub_runtime_op_num!(ConstantPad2dForward);
impl_stub_runtime_op_num!(ConstantPad3dForward);
impl_stub_runtime_op_num!(ReflectionPad1dForward);
impl_stub_runtime_op_num!(ReflectionPad2dForward);
impl_stub_runtime_op_num!(ReflectionPad3dForward);
impl_stub_runtime_op_num!(ReplicationPad1dForward);
impl_stub_runtime_op_num!(ReplicationPad2dForward);
impl_stub_runtime_op_num!(ReplicationPad3dForward);
impl_stub_runtime_op_num!(CircularPad1dForward);
impl_stub_runtime_op_num!(CircularPad2dForward);
impl_stub_runtime_op_num!(CircularPad3dForward);
impl_stub_runtime_op_float!(EluForward);
impl_stub_runtime_op_float!(SeluForward);
impl_stub_runtime_op_float!(CeluForward);
impl_stub_runtime_op_float!(MishForward);
impl_stub_runtime_op_float!(HardtanhForward);
impl_stub_runtime_op_float!(Relu6Forward);
impl_stub_runtime_op_float!(HardsigmoidForward);
impl_stub_runtime_op_float!(HardswishForward);
impl_stub_runtime_op_float!(HardshrinkForward);
impl_stub_runtime_op_float!(LeakyReluForward);
impl_stub_runtime_op_float!(ThresholdForward);
impl_stub_runtime_op_float!(SoftsignForward);
impl_stub_runtime_op_float!(SoftshrinkForward);
impl_stub_runtime_op_float!(SoftplusForward);
impl_stub_runtime_op_float!(LogsigmoidForward);
impl_stub_runtime_op_float!(TanhForward);
impl_stub_runtime_op_float!(TanhshrinkForward);
struct InputRuntimeOp;
impl RuntimeOp for InputRuntimeOp {
fn n_activation_inputs(&self) -> usize {
0
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
Vec::new()
}
fn pack_args(
&self,
_: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
_: teeny_core::model::RawPtr,
_: &[usize],
_: i32,
_: &mut dyn teeny_core::device::program::ArgVisitor,
) {
}
fn block(&self) -> [u32; 3] {
[1, 1, 1]
}
fn grid(&self, _: &[usize]) -> [u32; 3] {
[0, 0, 0]
}
}
#[derive(Debug, Default)]
pub struct TritonLowering {
sm_count: Option<u32>,
}
impl TritonLowering {
pub fn new() -> Self {
Self::default()
}
pub fn with_sm_count(mut self, sm_count: Option<u32>) -> Self {
self.sm_count = sm_count;
self
}
}
fn pick_adaptive_block_n(
tiled_dim: usize,
fixed_blocks: usize,
target_blocks: u32,
candidates: &[i32],
) -> i32 {
for &c in candidates {
let n_tiles = tiled_dim.div_ceil(c.max(1) as usize);
if (fixed_blocks * n_tiles) as u64 >= target_blocks as u64 {
return c;
}
}
*candidates.last().expect("candidates must be non-empty")
}
#[cfg(test)]
mod pick_adaptive_block_n_tests {
use super::pick_adaptive_block_n;
#[test]
fn keeps_largest_candidate_when_already_enough_blocks() {
let picked = pick_adaptive_block_n(256, 400, 512, &[16, 8, 4]);
assert_eq!(picked, 16);
}
#[test]
fn shrinks_for_occupancy_starved_shapes() {
let picked = pick_adaptive_block_n(256, 10, 512, &[16, 8, 4]);
assert_eq!(picked, 4);
}
#[test]
fn falls_back_to_smallest_candidate_when_target_unreachable() {
let picked = pick_adaptive_block_n(16, 1, 1_000_000, &[16, 8, 4]);
assert_eq!(picked, 4);
}
}
fn pick_gemm_tile_sizes(m: Option<usize>, n: usize, k: usize) -> (i32, i32, i32) {
let m = m.unwrap_or(64);
let block_k = if k >= 128 {
32
} else if k >= 32 {
16
} else {
8
};
let (block_m, block_n) = match (m, n) {
(m, n) if m >= 128 && n >= 128 => (128, 128),
(m, _) if m >= 128 => (128, 64),
(_, n) if n >= 128 => (64, 128),
_ => (64, 64),
};
(block_m, block_n, block_k)
}
#[cfg(test)]
mod pick_gemm_tile_sizes_tests {
use super::pick_gemm_tile_sizes;
#[test]
fn small_shape_gets_smallest_tiles() {
assert_eq!(pick_gemm_tile_sizes(Some(64), 64, 16), (64, 64, 8));
}
#[test]
fn large_m_and_n_get_the_largest_tile() {
assert_eq!(pick_gemm_tile_sizes(Some(512), 256, 256), (128, 128, 32));
}
#[test]
fn large_m_only_widens_block_m_not_block_n() {
assert_eq!(pick_gemm_tile_sizes(Some(512), 32, 64), (128, 64, 16));
}
#[test]
fn large_n_only_widens_block_n_not_block_m() {
assert_eq!(pick_gemm_tile_sizes(Some(32), 512, 64), (64, 128, 16));
}
#[test]
fn unknown_dynamic_batch_treated_as_small() {
assert_eq!(
pick_gemm_tile_sizes(None, 64, 16),
pick_gemm_tile_sizes(Some(64), 64, 16)
);
}
#[test]
fn block_k_never_drops_below_the_tensor_core_minimum() {
let (_, _, block_k) = pick_gemm_tile_sizes(Some(64), 64, 1);
assert!(block_k >= 8);
}
#[test]
fn block_k_grows_with_k() {
assert_eq!(pick_gemm_tile_sizes(Some(64), 64, 8).2, 8);
assert_eq!(pick_gemm_tile_sizes(Some(64), 64, 32).2, 16);
assert_eq!(pick_gemm_tile_sizes(Some(64), 64, 128).2, 32);
}
}
impl TritonLowering {
pub fn lower_with_mapping(
&self,
graph: &Graph,
mode: LoweringMode,
) -> Result<(Dag<Box<dyn ExecutableOp>>, Vec<usize>)> {
let _ = mode; let node_indexes = graph.topological_sort();
let mut dag: Dag<Box<dyn ExecutableOp>> = Dag::new();
let mut graph_to_dag = vec![0usize; graph.nodes.len()];
for node_index in node_indexes {
let node = &graph.nodes[node_index];
#[cfg(feature = "training")]
if mode == LoweringMode::Training
&& let Op::BatchNorm1d {
num_features,
eps,
momentum,
..
}
| Op::BatchNorm3d {
num_features,
eps,
momentum,
..
} = &node.op
{
let c = *num_features;
let eps_f32 = *eps as f32;
let momentum_f32 = *momentum as f32;
const BLOCK_N: i32 = 64;
let (stats_name, stats_src, stats_rop): (String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let k = BatchNormStatsForward::<f32>::new(BLOCK_N);
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(
BatchNormStatsRuntimeOp::<f32>::new(BLOCK_N, eps_f32, momentum_f32),
);
(k.name.to_string(), src, rop)
}
DtypeRepr::F64 => {
let k = BatchNormStatsForward::<f64>::new(BLOCK_N);
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(
BatchNormStatsRuntimeOp::<f64>::new(BLOCK_N, eps_f32, momentum_f32),
);
(k.name.to_string(), src, rop)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not a Float dtype for BatchNormStatsForward",
other
));
}
};
let stats_node = Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", stats_name),
name: stats_name,
kernel_source: stats_src,
shape: vec![Some(2 * c)],
dtype: node.dtype,
backward_kernel_source: String::new(),
backward_entry_point: String::new(),
runtime_op: stats_rop,
}) as Box<dyn ExecutableOp>;
let stats_dag_idx = dag.add_node(stats_node);
for &input_graph_idx in &node.inputs {
dag.add_edge(graph_to_dag[input_graph_idx], stats_dag_idx);
}
let (norm_name, norm_src, norm_bwd_src, norm_rop): (
String,
String,
String,
Arc<dyn RuntimeOp>,
) = match node.dtype {
DtypeRepr::F32 => {
let k = BatchNormNormalizeForward::<f32>::new(BLOCK_N);
let src = k.source.clone();
let rop = BatchNormNormalizeRuntimeOp::<f32>::new(BLOCK_N);
let bwd_src = rop.backward_source().to_string();
(
k.name.to_string(),
src,
bwd_src,
Arc::new(rop) as Arc<dyn RuntimeOp>,
)
}
DtypeRepr::F64 => {
let k = BatchNormNormalizeForward::<f64>::new(BLOCK_N);
let src = k.source.clone();
let rop = BatchNormNormalizeRuntimeOp::<f64>::new(BLOCK_N);
let bwd_src = rop.backward_source().to_string();
(
k.name.to_string(),
src,
bwd_src,
Arc::new(rop) as Arc<dyn RuntimeOp>,
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not a Float dtype for BatchNormNormalizeForward",
other
));
}
};
let norm_node = Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", norm_name),
name: norm_name,
kernel_source: norm_src,
shape: node.shape.clone(),
dtype: node.dtype,
backward_kernel_source: norm_bwd_src,
backward_entry_point: String::new(),
runtime_op: norm_rop,
}) as Box<dyn ExecutableOp>;
let norm_dag_idx = dag.add_node(norm_node);
for &input_graph_idx in &node.inputs {
dag.add_edge(graph_to_dag[input_graph_idx], norm_dag_idx);
}
dag.add_edge(stats_dag_idx, norm_dag_idx);
graph_to_dag[node_index] = norm_dag_idx;
continue;
}
if mode == LoweringMode::Inference
&& let Op::Conv2d {
has_bias: true,
kernel_h,
kernel_w,
stride_h,
stride_w,
padding_h,
padding_w,
groups,
..
} = &node.op
{
let (name, ks, rop): (String, String, Arc<dyn RuntimeOp>) = match node.dtype {
DtypeRepr::F32 => {
let k = Conv2dBiasForward::<f32>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
);
let nm = k.name.to_string();
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(k);
(nm, src, rop)
}
DtypeRepr::F64 => {
let k = Conv2dBiasForward::<f64>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
);
let nm = k.name.to_string();
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(k);
(nm, src, rop)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for Conv2dBiasForward",
other
));
}
};
let dag_idx = dag.add_node(Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
}) as Box<dyn ExecutableOp>);
for &input_graph_idx in &node.inputs {
dag.add_edge(graph_to_dag[input_graph_idx], dag_idx);
}
graph_to_dag[node_index] = dag_idx;
continue;
}
if let Op::Conv2d {
has_bias: true,
kernel_h,
kernel_w,
stride_h,
stride_w,
padding_h,
padding_w,
groups,
out_channels,
..
} = &node.op
{
const BIAS_BLOCK_HW: i32 = 128;
let (conv_name, conv_ks, conv_rop): (String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let k = Conv2dForward::<f32>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
);
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(Conv2dForward::<f32>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
));
(k.name.to_string(), src, rop)
}
DtypeRepr::F64 => {
let k = Conv2dForward::<f64>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
);
let src = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(Conv2dForward::<f64>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
));
(k.name.to_string(), src, rop)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for Conv2dForward",
other
));
}
};
#[cfg(feature = "training")]
let conv_bwd_ks = match node.dtype {
DtypeRepr::F32 => {
Conv2dBackward::<f32>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
)
.source
}
DtypeRepr::F64 => {
Conv2dBackward::<f64>::new(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16,
)
.source
}
_ => String::new(),
};
let conv_dag_idx = dag.add_node(Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", conv_name),
name: conv_name,
kernel_source: conv_ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: conv_bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: conv_rop,
}) as Box<dyn ExecutableOp>);
for &input_graph_idx in &node.inputs {
dag.add_edge(graph_to_dag[input_graph_idx], conv_dag_idx);
}
let (bias_name, bias_ks, bias_rop): (String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let r = NchwBiasAddRuntimeOp::<f32>::new(BIAS_BLOCK_HW);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = NchwBiasAddRuntimeOp::<f64>::new(BIAS_BLOCK_HW);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for NchwBiasAdd",
other
));
}
};
#[cfg(feature = "training")]
let bias_bwd_ks = match node.dtype {
DtypeRepr::F32 => NchwBiasAddRuntimeOp::<f32>::new(BIAS_BLOCK_HW)
.backward_source()
.to_string(),
DtypeRepr::F64 => NchwBiasAddRuntimeOp::<f64>::new(BIAS_BLOCK_HW)
.backward_source()
.to_string(),
_ => String::new(),
};
let biasadd_dag_idx = dag.add_node(Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", bias_name),
name: bias_name,
kernel_source: bias_ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bias_bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: bias_rop,
}) as Box<dyn ExecutableOp>);
dag.add_edge(conv_dag_idx, biasadd_dag_idx);
graph_to_dag[node_index] = biasadd_dag_idx;
let _ = (conv_dag_idx, out_channels);
continue;
}
let executable: Box<dyn ExecutableOp> = match &node.op {
Op::Input => Box::new(KernelExecutable {
name: "input".to_string(),
kernel_source: String::new(),
entry_point: String::new(),
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: Arc::new(InputRuntimeOp),
}),
Op::Linear { has_bias, .. } => {
make_num_kernel!(
LinearForward(*has_bias, 32, 64, 32, 8),
LinearBackward(*has_bias, 32, 64, 32, 8),
node
)
}
Op::Flatten => make_num_kernel!(FlattenForward(32, 256), node),
Op::BatchNorm1d { .. } | Op::BatchNorm3d { .. } => {
make_float_kernel!(BatchNormForwardInference(64), node)
}
Op::BatchNorm2d { eps, .. } => {
let eps_f32 = *eps as f32;
const BN2D_BLOCK_HW: i32 = 128;
let (name, ks, rop): (String, String, Arc<dyn RuntimeOp>) = match node.dtype {
DtypeRepr::F32 => {
let r = BatchNorm2dNchwInferenceRuntimeOp::<f32>::new(
BN2D_BLOCK_HW,
eps_f32,
);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = BatchNorm2dNchwInferenceRuntimeOp::<f64>::new(
BN2D_BLOCK_HW,
eps_f32,
);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not a Float dtype for BatchNorm2d",
other
));
}
};
#[cfg(feature = "training")]
let bwd_ks = match node.dtype {
DtypeRepr::F32 => BatchNorm2dNchwBackward::<f32>::new(BN2D_BLOCK_HW).source,
DtypeRepr::F64 => BatchNorm2dNchwBackward::<f64>::new(BN2D_BLOCK_HW).source,
_ => String::new(),
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::LayerNorm { eps, .. } => {
let eps_f32 = *eps as f32;
const LN_BLOCK_N: i32 = 1024;
let (name, ks, rop): (String, String, Arc<dyn RuntimeOp>) = match node.dtype {
DtypeRepr::F32 => {
let r =
LayerNormForwardInferenceRuntimeOp::<f32>::new(LN_BLOCK_N, eps_f32);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r =
LayerNormForwardInferenceRuntimeOp::<f64>::new(LN_BLOCK_N, eps_f32);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not a Float dtype for LayerNorm",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::RmsNorm { .. } => {
make_float_kernel!(RmsNormForward(1024), node)
}
Op::GroupNorm { .. } => {
make_float_kernel!(GroupNormForwardInference(256), node)
}
Op::InstanceNorm1d { .. }
| Op::InstanceNorm2d { .. }
| Op::InstanceNorm3d { .. } => {
make_float_kernel!(InstanceNormForwardInference(256), node)
}
Op::Conv1d {
kernel_l,
stride,
padding,
..
} => {
make_num_kernel!(
Conv1dForward(*kernel_l as i32, *stride as i32, *padding as i32, 32),
node
)
}
Op::Conv2d {
kernel_h,
kernel_w,
stride_h,
stride_w,
padding_h,
padding_w,
groups,
..
} => {
make_num_kernel!(
Conv2dForward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16
),
Conv2dBackward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*padding_h as i32,
*padding_w as i32,
*groups as i32,
16
),
node
)
}
Op::Conv3d {
kernel_d,
kernel_h,
kernel_w,
stride_d,
stride_h,
stride_w,
padding_d,
padding_h,
padding_w,
..
} => {
make_num_kernel!(
Conv3dForward(
*kernel_d as i32,
*kernel_h as i32,
*kernel_w as i32,
*stride_d as i32,
*stride_h as i32,
*stride_w as i32,
*padding_d as i32,
*padding_h as i32,
*padding_w as i32,
8
),
node
)
}
Op::Conv2dBnSilu {
kernel_h,
kernel_w,
stride_h,
stride_w,
padding_h,
padding_w,
groups,
in_channels,
..
} => {
if node.dtype != DtypeRepr::F32 {
return Err(anyhow::anyhow!(
"Conv2dBnSilu only supports f32 (got {:?})",
node.dtype
));
}
let kh = *kernel_h as i32;
let kw = *kernel_w as i32;
let sh = *stride_h as i32;
let sw = *stride_w as i32;
let ph = *padding_h as i32;
let pw = *padding_w as i32;
let g = *groups as i32;
let c_out = node.shape[1].unwrap_or(0);
let oh = node.shape[2].unwrap_or(1);
let ow = node.shape[3].unwrap_or(1);
let is_depthwise = g as usize == *in_channels;
let use_gemm = kh == 1
&& kw == 1
&& sh == 1
&& sw == 1
&& ph == 0
&& pw == 0
&& !is_depthwise
&& g == 1
&& c_out >= 32;
let use_tiled = !is_depthwise && g == 1 && c_out >= 16 && !use_gemm;
if use_gemm {
const GROUP_M: i32 = 8;
let m = oh * ow;
let (block_m, block_n_base, block_k) =
pick_gemm_tile_sizes(Some(m), c_out, *in_channels);
let block_n = match self.sm_count {
Some(sm_count) => {
let fixed_blocks = m.div_ceil(block_m as usize);
pick_adaptive_block_n(
c_out,
fixed_blocks,
4 * sm_count,
&[block_n_base, 16, 8],
)
}
None => block_n_base,
};
let k = Conv2dBnSiluGemmForward::new(block_m, block_n, block_k, GROUP_M);
let nm = k.name.to_string();
let ks = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(k);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
} else if use_tiled {
const BLOCK_OW: i32 = 16;
const BLOCK_N_TILE: i32 = 16;
let block_n_tile = match self.sm_count {
Some(sm_count) => {
let fixed_blocks = oh * ow.div_ceil(BLOCK_OW as usize);
pick_adaptive_block_n(
c_out,
fixed_blocks,
4 * sm_count,
&[BLOCK_N_TILE, 8, 4],
)
}
None => BLOCK_N_TILE,
};
let k = Conv2dBnSiluTiledForward::new(
kh,
kw,
sh,
sw,
ph,
pw,
BLOCK_OW,
block_n_tile,
);
let nm = k.name.to_string();
let ks = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(k);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
} else {
const BLOCK_OW: i32 = 16;
let k = Conv2dBnSiluForward::new(kh, kw, sh, sw, ph, pw, g, BLOCK_OW);
let nm = k.name.to_string();
let ks = k.source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(k);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
}
Op::AvgPool1d { kernel_l, stride } => {
make_num_kernel!(Avgpool1dForward(*kernel_l as i32, *stride as i32, 32), node)
}
Op::AvgPool2d {
kernel_h,
kernel_w,
stride_h,
stride_w,
} => {
make_num_kernel!(
Avgpool2dForward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
16
),
node
)
}
Op::AvgPool3d {
kernel_d,
kernel_h,
kernel_w,
stride_d,
stride_h,
stride_w,
} => {
make_num_kernel!(
Avgpool3dForward(
*kernel_d as i32,
*kernel_h as i32,
*kernel_w as i32,
*stride_d as i32,
*stride_h as i32,
*stride_w as i32,
8
),
node
)
}
Op::MaxPool1d { kernel_l, stride } => {
make_num_kernel!(Maxpool1dForward(*kernel_l as i32, *stride as i32, 32), node)
}
Op::MaxPool2d {
kernel_h,
kernel_w,
stride_h,
stride_w,
pad_h,
pad_w,
} => {
make_num_kernel!(
Maxpool2dForward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*pad_h as i32,
*pad_w as i32,
16
),
Maxpool2dBackward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
*pad_h as i32,
*pad_w as i32,
16
),
node
)
}
Op::MaxPool3d {
kernel_d,
kernel_h,
kernel_w,
stride_d,
stride_h,
stride_w,
} => {
make_num_kernel!(
Maxpool3dForward(
*kernel_d as i32,
*kernel_h as i32,
*kernel_w as i32,
*stride_d as i32,
*stride_h as i32,
*stride_w as i32,
8
),
node
)
}
Op::LpPool1d {
kernel_l, stride, ..
} => {
make_float_kernel!(Lppool1dForward(*kernel_l as i32, *stride as i32, 32), node)
}
Op::LpPool2d {
kernel_h,
kernel_w,
stride_h,
stride_w,
..
} => {
make_float_kernel!(
Lppool2dForward(
*kernel_h as i32,
*kernel_w as i32,
*stride_h as i32,
*stride_w as i32,
16
),
node
)
}
Op::LpPool3d {
kernel_d,
kernel_h,
kernel_w,
stride_d,
stride_h,
stride_w,
..
} => {
make_float_kernel!(
Lppool3dForward(
*kernel_d as i32,
*kernel_h as i32,
*kernel_w as i32,
*stride_d as i32,
*stride_h as i32,
*stride_w as i32,
8
),
node
)
}
Op::ConstantPad1d {
pad_left,
pad_right,
..
} => {
make_num_kernel!(
ConstantPad1dForward(*pad_left as i32, *pad_right as i32, 32),
node
)
}
Op::ConstantPad2d {
pad_l,
pad_r,
pad_t,
pad_b,
..
} => {
make_num_kernel!(
ConstantPad2dForward(
*pad_t as i32,
*pad_b as i32,
*pad_l as i32,
*pad_r as i32,
16
),
node
)
}
Op::ConstantPad3d {
pad_d1,
pad_d2,
pad_h1,
pad_h2,
pad_w1,
pad_w2,
..
} => {
make_num_kernel!(
ConstantPad3dForward(
*pad_d1 as i32,
*pad_d2 as i32,
*pad_h1 as i32,
*pad_h2 as i32,
*pad_w1 as i32,
*pad_w2 as i32,
8
),
node
)
}
Op::ReflectionPad1d {
pad_left,
pad_right,
} => {
make_num_kernel!(
ReflectionPad1dForward(*pad_left as i32, *pad_right as i32, 32),
node
)
}
Op::ReflectionPad2d {
pad_l,
pad_r,
pad_t,
pad_b,
} => {
make_num_kernel!(
ReflectionPad2dForward(
*pad_t as i32,
*pad_b as i32,
*pad_l as i32,
*pad_r as i32,
16
),
node
)
}
Op::ReflectionPad3d {
pad_d1,
pad_d2,
pad_h1,
pad_h2,
pad_w1,
pad_w2,
} => {
make_num_kernel!(
ReflectionPad3dForward(
*pad_d1 as i32,
*pad_d2 as i32,
*pad_h1 as i32,
*pad_h2 as i32,
*pad_w1 as i32,
*pad_w2 as i32,
8
),
node
)
}
Op::ReplicationPad1d {
pad_left,
pad_right,
} => {
make_num_kernel!(
ReplicationPad1dForward(*pad_left as i32, *pad_right as i32, 32),
node
)
}
Op::ReplicationPad2d {
pad_l,
pad_r,
pad_t,
pad_b,
} => {
make_num_kernel!(
ReplicationPad2dForward(
*pad_t as i32,
*pad_b as i32,
*pad_l as i32,
*pad_r as i32,
16
),
node
)
}
Op::ReplicationPad3d {
pad_d1,
pad_d2,
pad_h1,
pad_h2,
pad_w1,
pad_w2,
} => {
make_num_kernel!(
ReplicationPad3dForward(
*pad_d1 as i32,
*pad_d2 as i32,
*pad_h1 as i32,
*pad_h2 as i32,
*pad_w1 as i32,
*pad_w2 as i32,
8
),
node
)
}
Op::CircularPad1d {
pad_left,
pad_right,
} => {
make_num_kernel!(
CircularPad1dForward(*pad_left as i32, *pad_right as i32, 32),
node
)
}
Op::CircularPad2d {
pad_l,
pad_r,
pad_t,
pad_b,
} => {
make_num_kernel!(
CircularPad2dForward(
*pad_t as i32,
*pad_b as i32,
*pad_l as i32,
*pad_r as i32,
16
),
node
)
}
Op::CircularPad3d {
pad_d1,
pad_d2,
pad_h1,
pad_h2,
pad_w1,
pad_w2,
} => {
make_num_kernel!(
CircularPad3dForward(
*pad_d1 as i32,
*pad_d2 as i32,
*pad_h1 as i32,
*pad_h2 as i32,
*pad_w1 as i32,
*pad_w2 as i32,
8
),
node
)
}
Op::Relu => make_num_kernel!(ReluForward(1024), ReluBackward(1024), node),
Op::Elu { .. } => exec_from(
node.shape.clone(),
node.dtype,
EluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Selu => exec_from(
node.shape.clone(),
node.dtype,
SeluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Celu { .. } => exec_from(
node.shape.clone(),
node.dtype,
CeluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Gelu => exec_from(
node.shape.clone(),
node.dtype,
GeluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Mish => exec_from(
node.shape.clone(),
node.dtype,
MishForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Hardtanh { .. } => exec_from(
node.shape.clone(),
node.dtype,
HardtanhForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Relu6 => exec_from(
node.shape.clone(),
node.dtype,
Relu6ForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Hardsigmoid => exec_from(
node.shape.clone(),
node.dtype,
HardsigmoidForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Hardswish => exec_from(
node.shape.clone(),
node.dtype,
HardswishForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Hardshrink { .. } => exec_from(
node.shape.clone(),
node.dtype,
HardshrinkForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::LeakyRelu { .. } => exec_from(
node.shape.clone(),
node.dtype,
LeakyReluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Threshold { .. } => exec_from(
node.shape.clone(),
node.dtype,
ThresholdForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Softsign => exec_from(
node.shape.clone(),
node.dtype,
SoftsignForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Softshrink { .. } => exec_from(
node.shape.clone(),
node.dtype,
SoftshrinkForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Softplus { .. } => exec_from(
node.shape.clone(),
node.dtype,
SoftplusForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Sigmoid => exec_from(
node.shape.clone(),
node.dtype,
SigmoidForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Silu => exec_from(
node.shape.clone(),
node.dtype,
SiluForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Logsigmoid => exec_from(
node.shape.clone(),
node.dtype,
LogsigmoidForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Tanh => exec_from(
node.shape.clone(),
node.dtype,
TanhForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Tanhshrink => exec_from(
node.shape.clone(),
node.dtype,
TanhshrinkForwardDispatch::dispatch(node.dtype, 1024)?,
),
Op::Softmax { .. } => {
let n_cols = node.shape.last().and_then(|d| *d).unwrap_or(1024);
let block_size = n_cols.next_power_of_two() as i32;
make_float_kernel!(SoftmaxForward(block_size), node)
}
Op::UpsampleNearest2d { scale_h, scale_w } => {
make_num_kernel!(
UpsampleNearest2dForward(*scale_h as i32, *scale_w as i32, 16),
UpsampleNearest2dBackward(*scale_h as i32, *scale_w as i32, 16),
node
)
}
Op::ChannelCat { .. } => {
let n_inputs = node.inputs.len();
let (name, fwd_src, bwd_src, rop): (
String,
String,
String,
Arc<dyn RuntimeOp>,
) = match node.dtype {
DtypeRepr::F32 => {
let r = ChannelCatRuntimeOp::<f32>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = ChannelCatRuntimeOp::<f64>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I8 => {
let r = ChannelCatRuntimeOp::<i8>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I16 => {
let r = ChannelCatRuntimeOp::<i16>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I32 => {
let r = ChannelCatRuntimeOp::<i32>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I64 => {
let r = ChannelCatRuntimeOp::<i64>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U8 => {
let r = ChannelCatRuntimeOp::<u8>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U16 => {
let r = ChannelCatRuntimeOp::<u16>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U32 => {
let r = ChannelCatRuntimeOp::<u32>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U64 => {
let r = ChannelCatRuntimeOp::<u64>::new(128, n_inputs);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for ChannelCat",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::ChannelChunk {
chunk_c,
chunk_offset,
..
} => {
let chunk_c = *chunk_c;
let chunk_offset = *chunk_offset;
let (name, fwd_src, bwd_src, rop): (
String,
String,
String,
Arc<dyn RuntimeOp>,
) = match node.dtype {
DtypeRepr::F32 => {
let r = ChannelChunkRuntimeOp::<f32>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = ChannelChunkRuntimeOp::<f64>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I8 => {
let r = ChannelChunkRuntimeOp::<i8>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I16 => {
let r = ChannelChunkRuntimeOp::<i16>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I32 => {
let r = ChannelChunkRuntimeOp::<i32>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::I64 => {
let r = ChannelChunkRuntimeOp::<i64>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U8 => {
let r = ChannelChunkRuntimeOp::<u8>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U16 => {
let r = ChannelChunkRuntimeOp::<u16>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U32 => {
let r = ChannelChunkRuntimeOp::<u32>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::U64 => {
let r = ChannelChunkRuntimeOp::<u64>::new(128, chunk_c, chunk_offset);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for ChannelChunk",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::ChannelBiasAdd { c } => {
let c = *c;
let (name, fwd_src, bwd_src, rop): (
String,
String,
String,
Arc<dyn RuntimeOp>,
) = match node.dtype {
DtypeRepr::F32 => {
let r = ChannelBiasAddRuntimeOp::<f32>::new(128, c);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = ChannelBiasAddRuntimeOp::<f64>::new(128, c);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for ChannelBiasAdd",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::Add => {
make_num_kernel!(ElemwiseAddForward(128), ElemwiseAddBackward(128), node)
}
Op::Abs => {
make_num_kernel!(ElemwiseAbsForward(1024), ElemwiseAbsBackward(1024), node)
}
Op::Neg => {
make_num_kernel!(ElemwiseNegForward(1024), ElemwiseNegBackward(1024), node)
}
Op::Sign => make_num_kernel!(ElemwiseSignForward(1024), node),
Op::IsNaN => make_float_kernel!(ElemwiseIsnanForward(1024), node),
Op::Ceil => make_float_kernel!(ElemwiseCeilForward(1024), node),
Op::Floor => make_float_kernel!(ElemwiseFloorForward(1024), node),
Op::Sqrt => {
make_float_kernel!(ElemwiseSqrtForward(1024), ElemwiseSqrtBackward(1024), node)
}
Op::Reciprocal => make_float_kernel!(
ElemwiseReciprocalForward(1024),
ElemwiseReciprocalBackward(1024),
node
),
Op::Exp => {
make_float_kernel!(ElemwiseExpForward(1024), ElemwiseExpBackward(1024), node)
}
Op::Log => {
make_float_kernel!(ElemwiseLogForward(1024), ElemwiseLogBackward(1024), node)
}
Op::Erf => {
make_float_kernel!(ElemwiseErfForward(1024), ElemwiseErfBackward(1024), node)
}
Op::Sin => {
make_float_kernel!(ElemwiseSinForward(1024), ElemwiseSinBackward(1024), node)
}
Op::Cos => {
make_float_kernel!(ElemwiseCosForward(1024), ElemwiseCosBackward(1024), node)
}
Op::Tan => {
make_float_kernel!(ElemwiseTanForward(1024), ElemwiseTanBackward(1024), node)
}
Op::Asin => {
make_float_kernel!(ElemwiseAsinForward(1024), ElemwiseAsinBackward(1024), node)
}
Op::Acos => {
make_float_kernel!(ElemwiseAcosForward(1024), ElemwiseAcosBackward(1024), node)
}
Op::Atan => {
make_float_kernel!(ElemwiseAtanForward(1024), ElemwiseAtanBackward(1024), node)
}
Op::Sinh => {
make_float_kernel!(ElemwiseSinhForward(1024), ElemwiseSinhBackward(1024), node)
}
Op::Cosh => {
make_float_kernel!(ElemwiseCoshForward(1024), ElemwiseCoshBackward(1024), node)
}
Op::Asinh => make_float_kernel!(
ElemwiseAsinhForward(1024),
ElemwiseAsinhBackward(1024),
node
),
Op::Acosh => make_float_kernel!(
ElemwiseAcoshForward(1024),
ElemwiseAcoshBackward(1024),
node
),
Op::Atanh => make_float_kernel!(
ElemwiseAtanhForward(1024),
ElemwiseAtanhBackward(1024),
node
),
Op::Round => {
return Err(anyhow::anyhow!(
"TODO: Op::Round — implement rounding kernel"
));
}
Op::Mul => {
make_num_kernel!(ElemwiseMulForward(1024), ElemwiseMulBackward(1024), node)
}
Op::Sub => {
make_num_kernel!(ElemwiseSubForward(1024), ElemwiseSubBackward(1024), node)
}
Op::Div => {
make_float_kernel!(ElemwiseDivForward(1024), ElemwiseDivBackward(1024), node)
}
Op::Pow => {
make_float_kernel!(ElemwisePowForward(1024), ElemwisePowBackward(1024), node)
}
Op::Mod { .. } => make_float_kernel!(ElemwiseFmodForward(1024), node),
Op::ElemMin => {
make_num_kernel!(ElemwiseMinForward(1024), ElemwiseMinBackward(1024), node)
}
Op::ElemMax => {
make_num_kernel!(ElemwiseMaxForward(1024), ElemwiseMaxBackward(1024), node)
}
Op::ElemMean => {
make_float_kernel!(ElemwiseMeanForward(1024), ElemwiseMeanBackward(1024), node)
}
Op::ElemSum => {
make_num_kernel!(ElemwiseSumForward(1024), ElemwiseSumBackward(1024), node)
}
Op::Equal => make_num_kernel!(ElemwiseEqualForward(1024), node),
Op::Greater => make_num_kernel!(ElemwiseGreaterForward(1024), node),
Op::GreaterOrEqual => make_num_kernel!(ElemwiseGreaterEqualForward(1024), node),
Op::Less => make_num_kernel!(ElemwiseLessForward(1024), node),
Op::LessOrEqual => make_num_kernel!(ElemwiseLessEqualForward(1024), node),
Op::Where => make_float_kernel!(
ElemwiseWhereForward(1024),
ElemwiseWhereBackward(1024),
node
),
Op::Clip => {
let (name, ks, bwd_ks, rop): (String, String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let r = ClipRuntimeOp::<f32>::new(
1024,
f32::NEG_INFINITY,
f32::INFINITY,
);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = ClipRuntimeOp::<f64>::new(
1024,
f32::NEG_INFINITY,
f32::INFINITY,
);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for Clip",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::ReduceSum { .. } => make_num_kernel!(ReduceSumForward(1024), node),
Op::ReduceMean { .. } => make_float_kernel!(ReduceMeanForward(1024), node),
Op::ReduceMax { .. } => make_num_kernel!(ReduceMaxForward(1024), node),
Op::ReduceMin { .. } => make_num_kernel!(ReduceMinForward(1024), node),
Op::ReduceProd { .. } => make_float_kernel!(ReduceProdForward(1024), node),
Op::ReduceL1 { .. } => make_num_kernel!(ReduceL1Forward(1024), node),
Op::ReduceL2 { .. } => make_float_kernel!(ReduceL2Forward(1024), node),
Op::ReduceLogSum { .. } => make_float_kernel!(ReduceLogSumForward(1024), node),
Op::ReduceLogSumExp { .. } => {
make_float_kernel!(ReduceLogSumExpForward(1024), node)
}
Op::ReduceSumSquare { .. } => make_num_kernel!(ReduceSumSquareForward(1024), node),
Op::CumSum { .. } => make_num_kernel!(CumSumForward(1024), node),
Op::CumProd { .. } => make_num_kernel!(CumProdForward(1024), node),
Op::GlobalAvgPool => make_float_kernel!(GlobalAvgPoolForward(1024), node),
Op::GlobalMaxPool => make_float_kernel!(GlobalMaxPoolForward(1024), node),
Op::ArgMax { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::ArgMax — I32Tensor output requires a custom kernel"
));
}
Op::ArgMin { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::ArgMin — I32Tensor output requires a custom kernel"
));
}
Op::Swish => {
let fwd = SwishForward::new(1024);
let nm = fwd.name.to_string();
let fwd_src = fwd.source.clone();
#[cfg(feature = "training")]
let bwd_src = SwishBackward::new(1024).source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(fwd);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::PRelu => {
let fwd = PreluForward::new(1024);
let nm = fwd.name.to_string();
let fwd_src = fwd.source.clone();
#[cfg(feature = "training")]
let bwd_src = crate::nn::activation::extra::PreluBackward::new(1024)
.source
.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(fwd);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::LogSoftmax { .. } => {
let n_cols = node.shape.last().and_then(|d| *d).unwrap_or(1024);
let block_size = n_cols.next_power_of_two() as i32;
let fwd = LogSoftmaxForward::new(block_size);
let nm = fwd.name.to_string();
let fwd_src = fwd.source.clone();
#[cfg(feature = "training")]
let bwd_src = LogSoftmaxBackward::new(block_size).source.clone();
let rop: Arc<dyn RuntimeOp> = Arc::new(fwd);
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", nm),
name: nm,
kernel_source: fwd_src,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_src,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::ThresholdedRelu { alpha } => {
let alpha_f = *alpha as f32;
let (name, ks, bwd_ks, rop): (String, String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let r = ThresholdedReluRuntimeOp::new(1024, alpha_f);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for ThresholdedRelu",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::Shrink { lambd, bias } => {
let lambd_f = *lambd as f32;
let bias_f = *bias as f32;
let (name, ks, bwd_ks, rop): (String, String, String, Arc<dyn RuntimeOp>) =
match node.dtype {
DtypeRepr::F32 => {
let r = ShrinkRuntimeOp::new(1024, lambd_f, bias_f);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
r.backward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not supported for Shrink",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: bwd_ks,
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::MatMul | Op::Gemm { .. } => {
let m = node.shape.first().copied().flatten();
let n = node.shape.last().copied().flatten().unwrap_or(0);
let k = node
.inputs
.first()
.and_then(|&i| graph.nodes[i].shape.last().copied().flatten())
.unwrap_or(0);
let (block_m, block_n, block_k) = pick_gemm_tile_sizes(m, n, k);
let (name, ks, rop): (String, String, Arc<dyn RuntimeOp>) = match node.dtype {
DtypeRepr::F32 => {
let r = MatMulRuntimeOp::<f32>::new(block_m, block_n, block_k);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
DtypeRepr::F64 => {
let r = MatMulRuntimeOp::<f64>::new(block_m, block_n, block_k);
(
r.kernel_name().to_string(),
r.forward_source().to_string(),
Arc::new(r),
)
}
other => {
return Err(anyhow::anyhow!(
"{:?} is not a Float dtype for MatMul",
other
));
}
};
Box::new(KernelExecutable {
entry_point: format!("{}_entry_point", name),
name,
kernel_source: ks,
shape: node.shape.clone(),
dtype: node.dtype,
#[cfg(feature = "training")]
backward_kernel_source: String::new(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
runtime_op: rop,
})
}
Op::Lstm { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Lstm — multi-step recurrent; implement as a custom loop kernel"
));
}
Op::Gru { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Gru — multi-step recurrent; implement as a custom loop kernel"
));
}
Op::Rnn { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Rnn — multi-step recurrent; implement as a custom loop kernel"
));
}
Op::RotaryEmbedding => {
return Err(anyhow::anyhow!(
"TODO: Op::RotaryEmbedding — implement RoPE kernel"
));
}
Op::MultiHeadAttention { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::MultiHeadAttention — use flash attention or a custom MHA kernel"
));
}
Op::FlexAttention { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::FlexAttention — implement a custom attention kernel supporting an arbitrary score-modification function"
));
}
Op::LinearAttention { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::LinearAttention — implement a linear-attention/gated-delta-rule kernel"
));
}
Op::CausalConvWithState { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::CausalConvWithState — implement a stateful causal-conv kernel"
));
}
Op::Reshape => {
return Err(anyhow::anyhow!(
"TODO: Op::Reshape — implement as a strided view or copy kernel"
));
}
Op::Transpose { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Transpose — implement as a permuted-copy kernel"
));
}
Op::Squeeze { .. } | Op::Unsqueeze { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Squeeze/Unsqueeze — implement as a zero-copy view"
));
}
Op::Concat { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Concat — implement as a multi-input copy kernel"
));
}
Op::Split { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Split — implement as a multi-output slice kernel"
));
}
Op::Slice => {
return Err(anyhow::anyhow!(
"TODO: Op::Slice — implement as a strided-copy kernel"
));
}
Op::Gather { .. } | Op::GatherElements { .. } | Op::GatherND { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Gather — implement as an index-gather kernel"
));
}
Op::ScatterElements { .. }
| Op::ScatterND
| Op::Scatter { .. }
| Op::TensorScatter => {
return Err(anyhow::anyhow!(
"TODO: Op::Scatter — implement as an index-scatter kernel"
));
}
Op::Tile => {
return Err(anyhow::anyhow!(
"TODO: Op::Tile — implement as a tiled-copy kernel"
));
}
Op::Expand => {
return Err(anyhow::anyhow!(
"TODO: Op::Expand — implement as a broadcast-copy kernel"
));
}
Op::ShapeOf { .. } | Op::SizeOf => {
return Err(anyhow::anyhow!(
"TODO: Op::ShapeOf/SizeOf — output is metadata, not tensor data"
));
}
Op::Compress { .. } | Op::NonZero => {
return Err(anyhow::anyhow!(
"TODO: Op::Compress/NonZero — variable-output ops require stream compaction"
));
}
Op::Range => {
return Err(anyhow::anyhow!(
"TODO: Op::Range — implement as a fill/arange kernel"
));
}
Op::Constant { .. } | Op::ConstantOfShape { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Constant — inline constant; should be materialised before lowering"
));
}
Op::Trilu { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Trilu — implement as a triangular mask kernel"
));
}
Op::Pad { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Pad — implement as a generic N-D padding kernel"
));
}
Op::ReverseSequence { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::ReverseSequence — implement as a scatter-copy kernel"
));
}
Op::Einsum { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Einsum — parse equation and emit a fused contraction kernel"
));
}
Op::Det => {
return Err(anyhow::anyhow!(
"TODO: Op::Det — implement via LU decomposition"
));
}
Op::QLinearMatMul
| Op::MatMulInteger
| Op::ConvInteger { .. }
| Op::QLinearConv { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Q* quantised matmul/conv — implement quantised compute kernels"
));
}
Op::ConvTranspose { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::ConvTranspose — implement transposed (gradient) convolution kernel"
));
}
Op::DeformConv { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::DeformConv — implement deformable convolution kernel"
));
}
Op::Col2Im { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Col2Im — implement col2im (fold) kernel"
));
}
Op::Resize { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Resize — implement nearest/bilinear/bicubic resize kernels"
));
}
Op::GridSample { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::GridSample — implement bilinear grid sample kernel"
));
}
Op::SpaceToDepth { .. } | Op::DepthToSpace { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::SpaceToDepth/DepthToSpace — implement pixel shuffle kernel"
));
}
Op::RoiAlign { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::RoiAlign — implement RoI-align pooling kernel"
));
}
Op::AffineGrid { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::AffineGrid — implement affine grid generator kernel"
));
}
Op::MaxUnpool { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::MaxUnpool — implement max-unpool (scatter with saved indices) kernel"
));
}
Op::CenterCropPad { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::CenterCropPad — implement center-crop-pad kernel"
));
}
Op::NonMaxSuppression { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::NonMaxSuppression — implement NMS kernel"
));
}
Op::TopK { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::TopK — implement radix sort / parallel selection kernel"
));
}
Op::Unique { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Unique — implement stream-compaction unique kernel"
));
}
Op::EyeLike { .. }
| Op::OneHot { .. }
| Op::Bernoulli { .. }
| Op::RandomUniformLike { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::EyeLike/OneHot/Bernoulli/RandomUniformLike — implement generation kernels"
));
}
Op::And | Op::Or | Op::Xor => {
return Err(anyhow::anyhow!(
"TODO: Op::And/Or/Xor — implement boolean logical kernels"
));
}
Op::BitShift { .. }
| Op::BitwiseAnd
| Op::BitwiseOr
| Op::BitwiseXor
| Op::BitwiseNot
| Op::Not => {
return Err(anyhow::anyhow!(
"TODO: Op::Bitwise* — implement integer bitwise kernels"
));
}
Op::QuantizeLinear { .. }
| Op::DequantizeLinear { .. }
| Op::DynamicQuantizeLinear => {
return Err(anyhow::anyhow!(
"TODO: Op::Quantize/Dequantize — implement quantisation kernels"
));
}
Op::LRN { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::LRN — implement local response normalisation kernel"
));
}
Op::MeanVarianceNormalization { .. } | Op::LpNormalization { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::MvnNorm/LpNorm — implement normalisation kernels"
));
}
Op::Dft { .. }
| Op::Stft
| Op::MelWeightMatrix
| Op::HannWindow { .. }
| Op::BlackmanWindow { .. }
| Op::HammingWindow { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::DFT/STFT/Window — implement signal processing kernels"
));
}
Op::NegativeLogLikelihoodLoss { .. } | Op::SoftmaxCrossEntropyLoss { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::NllLoss/SoftmaxCELoss — implement loss kernels"
));
}
Op::SequenceAt
| Op::SequenceConstruct
| Op::SequenceEmpty
| Op::SequenceErase
| Op::SequenceInsert
| Op::SequenceLength
| Op::SequenceMap
| Op::SplitToSequence { .. }
| Op::ConcatFromSequence { .. }
| Op::OptionalGetElement
| Op::OptionalHasElement
| Op::Loop
| Op::Scan { .. }
| Op::If => {
return Err(anyhow::anyhow!(
"TODO: Op::Sequence/Control-flow — not lowerable to single Triton kernels"
));
}
Op::Adagrad | Op::Adam | Op::Momentum | Op::Gradient => {
return Err(anyhow::anyhow!(
"TODO: Op::OnnxOptimizer — use teenygrad's own optimizer kernels instead"
));
}
Op::StringNormalizer
| Op::RegexFullMatch { .. }
| Op::StringConcat
| Op::StringSplit
| Op::TfIdfVectorizer
| Op::LabelEncoder
| Op::ArrayFeatureExtractor
| Op::Binarizer { .. }
| Op::TreeEnsemble
| Op::ImageDecoder => {
return Err(anyhow::anyhow!(
"TODO: Op::String/ClassicalML — not GPU-lowerable"
));
}
Op::CastLike | Op::BitCast { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::CastLike/BitCast — implement dtype-cast kernels"
));
}
Op::Cast { to } => {
return Err(anyhow::anyhow!(
"TODO: Op::Cast to {:?} — implement dtype-cast kernel",
to
));
}
Op::Identity => {
return Err(anyhow::anyhow!(
"TODO: Op::Identity — implement zero-copy pass-through kernel"
));
}
Op::Dropout { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Dropout — implement inference pass-through / training dropout kernel"
));
}
Op::IsInf { .. } => {
return Err(anyhow::anyhow!("TODO: Op::IsInf — implement isinf kernel"));
}
Op::Hardmax { .. } => {
return Err(anyhow::anyhow!(
"TODO: Op::Hardmax — implement argmax + one-hot kernel"
));
}
Op::Attention {
c,
num_heads,
key_dim,
} => {
let _ = (c, num_heads, key_dim);
return Err(anyhow::anyhow!(
"Op::Attention reached the match arm — this should not happen"
));
}
Op::Custom { data } => match data.0.lower() {
Some((name, kernel_source, entry_point, runtime_op)) => {
Box::new(KernelExecutable {
name,
kernel_source,
entry_point,
shape: node.shape.clone(),
dtype: node.dtype,
runtime_op,
#[cfg(feature = "training")]
backward_kernel_source: data.0.lower_backward_source(),
#[cfg(feature = "training")]
backward_entry_point: String::new(),
})
}
None => {
return Err(anyhow::anyhow!(
"custom op '{}' is not handled — implement CustomOp::lower()",
data.name()
));
}
},
Op::Fused { members } => {
return Err(anyhow::anyhow!(
"Op::Fused lowering is not implemented yet ({} member op(s)): \
concatenate each member's kernel source and synthesize an \
entry point that runs them in sequence (see spinorml-1fj.1)",
members.len()
));
}
};
let dag_idx = dag.add_node(executable);
graph_to_dag[node_index] = dag_idx;
for &input_graph_idx in &node.inputs {
dag.add_edge(graph_to_dag[input_graph_idx], dag_idx);
}
}
Ok((dag, graph_to_dag))
}
}
impl<'a> Lowering<'a> for TritonLowering {
fn lower(&self, graph: &Graph, mode: LoweringMode) -> Result<Dag<Box<dyn ExecutableOp>>> {
TritonLowering::lower_with_mapping(self, graph, mode).map(|(dag, _)| dag)
}
fn lower_with_mapping(
&self,
graph: &Graph,
mode: LoweringMode,
) -> Result<(Dag<Box<dyn ExecutableOp>>, Vec<usize>)> {
TritonLowering::lower_with_mapping(self, graph, mode)
}
fn extra_dag_names(&self, graph: &Graph, graph_to_dag: &[usize]) -> Vec<(usize, String)> {
let mut extra = Vec::new();
for (graph_idx, node) in graph.nodes.iter().enumerate() {
if let Op::Conv2d { has_bias: true, .. } = &node.op
&& let Some(name) = graph.names.get(&graph_idx)
{
let biasadd_dag_idx = graph_to_dag[graph_idx];
if biasadd_dag_idx > 0 {
extra.push((biasadd_dag_idx - 1, name.clone()));
}
}
}
extra
}
}