pub use crate::backend::{Dev, DeviceInfo};
pub use custom::{Acc, CompiledKernel, LocalPartition, Partition};
pub use ops::{BOp, MMADType, MMADims, MMALayout, OpId, ParamKind, TileDim};
pub(crate) use ops::{MoveOp, Op, OpLinked, RangeKind, UOp};
use crate::{DType, Map, Set, dtype::Constant, shape::Dim, slab::Slab};
use nanoserde::{DeBin, SerBin};
use std::collections::BTreeMap;
use std::sync::Arc;
use std::{hash::BuildHasherDefault, hash::Hash};
mod algebraic;
pub mod autotune;
mod coarsen;
mod cost;
mod custom;
mod debug;
mod fold_constants;
mod fold_loops;
mod fuse;
mod instr_sched;
mod licm;
mod linearize;
mod local_reduce;
mod merge_loops;
mod mma;
mod ops;
mod pad_range;
mod predict_cost;
mod split_loops;
mod tenstorrent;
mod transforms;
mod unroll_loops;
mod vectorize;
mod verify;
pub(crate) const IDX_T: DType = DType::I64;
#[derive(Debug, Clone)]
pub struct Kernel {
pub(crate) ops: Slab<OpId, OpLinked>,
pub(crate) head: OpId,
pub(crate) tail: OpId,
pub(crate) dev: Dev,
pub(crate) dev_info: Option<Arc<DeviceInfo>>,
pub(crate) shape_cache: Map<OpId, Vec<OpId>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, SerBin, DeBin)]
pub enum MemScope {
Global,
Local,
Register,
Circular,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, SerBin, DeBin)]
pub enum MemLayout {
Scalar,
Vector(u16),
Tile {
x: u16,
y: u16,
stride: u32,
},
}
impl PartialEq for Kernel {
fn eq(&self, other: &Self) -> bool {
self.ops == other.ops && self.head == other.head && self.dev == other.dev
}
}
impl Eq for Kernel {}
impl SerBin for Kernel {
fn ser_bin(&self, output: &mut Vec<u8>) {
self.ops.ser_bin(output);
self.head.ser_bin(output);
self.tail.ser_bin(output);
}
}
impl Hash for Kernel {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.head.hash(state);
self.ops.hash(state);
self.dev.hash(state);
}
}
impl Kernel {
pub(crate) fn from_device_id(dev: Dev, dev_info: Option<Arc<DeviceInfo>>) -> Self {
Self { ops: Slab::new(), head: OpId::NULL, tail: OpId::NULL, dev, dev_info, shape_cache: Map::default() }
}
pub(crate) fn dev_info(&self) -> &DeviceInfo {
self.dev_info.as_ref().expect("kernel has no device bound (Dev::Auto placeholder)")
}
pub(crate) fn compute_dtypes_and_rcs(&self) -> (Map<OpId, (DType, MemLayout)>, Map<OpId, u32>) {
let mut rcs: Map<OpId, u32> = Map::with_capacity_and_hasher(self.ops.len().into(), BuildHasherDefault::new());
let mut dtypes: Map<OpId, (DType, MemLayout)> = Map::with_capacity_and_hasher(100, BuildHasherDefault::new());
let mut op_id = self.head;
for _ in 0..10_000 {
if op_id.is_null() {
break;
}
match self.ops[op_id].op {
Op::Move { .. } | Op::Reduce { .. } | Op::ReduceTile { .. } => {
unreachable!()
}
Op::Const(x) => {
dtypes.insert(op_id, (x.dtype(), MemLayout::Scalar));
}
Op::Param { dtype, .. } => {
dtypes.insert(op_id, (dtype, MemLayout::Scalar));
}
Op::Storage { dtype, .. } => {
dtypes.insert(op_id, (dtype, MemLayout::Scalar));
}
Op::Load { src, index, layout } => {
dtypes.insert(op_id, (dtypes[&src].0, layout));
*rcs.entry(index).or_insert(0) += 1;
}
Op::Store { dst, src: x, index, layout } => {
debug_assert_eq!(dtypes[&x].1, layout);
dtypes.insert(op_id, dtypes[&x]);
*rcs.entry(dst).or_insert(0) += 1;
*rcs.entry(x).or_insert(0) += 1;
*rcs.entry(index).or_insert(0) += 1;
}
Op::Cast { x, dtype } => {
dtypes.insert(op_id, (dtype, dtypes[&x].1));
*rcs.entry(x).or_insert(0) += 1;
}
Op::Bitcast { x, dtype } => {
dtypes.insert(op_id, (dtype, dtypes[&x].1));
*rcs.entry(x).or_insert(0) += 1;
}
Op::Unary { x, .. } => {
dtypes.insert(op_id, dtypes[&x]);
*rcs.entry(x).or_insert(0) += 1;
}
Op::BroadcastTile { x, .. } => {
dtypes.insert(op_id, dtypes[&x]);
*rcs.entry(x).or_insert(0) += 1;
}
Op::Binary { x, y, bop } => {
let dtype = if bop.returns_bool() {
(DType::Bool, dtypes[&x].1)
} else {
dtypes[&x]
};
dtypes.insert(op_id, dtype);
*rcs.entry(x).or_insert(0) += 1;
*rcs.entry(y).or_insert(0) += 1;
}
Op::Asm { ref ops, .. } => {
let dtype = dtypes[&ops[0]];
dtypes.insert(op_id, dtype);
for &x in ops.iter() {
*rcs.entry(x).or_insert(0) += 1;
}
}
Op::Stack { ref ops } => {
let dtype = dtypes[&ops[0]];
dtypes.insert(op_id, (dtype.0, MemLayout::Vector(ops.len().try_into().unwrap())));
for &x in ops.iter() {
*rcs.entry(x).or_insert(0) += 1;
}
}
Op::Index { vec, idx: _ } => {
let dtype = dtypes[&vec];
dtypes.insert(op_id, (dtype.0, MemLayout::Scalar));
*rcs.entry(vec).or_insert(0) += 1;
}
Op::Wmma { dims: _, layout: _, dtype, a, b, c } => {
let out_dtype = match dtype {
MMADType::f16_f16_f16_f32 => DType::F32,
MMADType::f16_f16_f16_f16 => DType::F16,
MMADType::s8_s8_s32_s32
| MMADType::s4_s4_s32_s32
| MMADType::b1_b1_s32_xor_popc
| MMADType::b1_b1_s32_and_popc => DType::I32,
};
dtypes.insert(op_id, (out_dtype, MemLayout::Vector(4)));
*rcs.entry(a).or_insert(0) += 1;
*rcs.entry(b).or_insert(0) += 1;
*rcs.entry(c).or_insert(0) += 1;
}
Op::MatmulTile { x, y, acc } => {
dtypes.insert(op_id, dtypes[&acc]);
*rcs.entry(x).or_insert(0) += 1;
*rcs.entry(y).or_insert(0) += 1;
*rcs.entry(acc).or_insert(0) += 1;
}
Op::TransposeTile { x } => {
dtypes.insert(op_id, dtypes[&x]);
*rcs.entry(x).or_insert(0) += 1;
}
Op::Mad { x, y, z } => {
dtypes.insert(op_id, dtypes[&x]);
*rcs.entry(x).or_insert(0) += 1;
*rcs.entry(y).or_insert(0) += 1;
*rcs.entry(z).or_insert(0) += 1;
}
Op::Range { kind, .. } => {
if let RangeKind::Group(len) = kind {
*rcs.entry(len).or_insert(0) += 1;
}
if let RangeKind::Warp(local_id) = kind {
*rcs.entry(local_id).or_insert(0) += 1;
}
dtypes.insert(op_id, (IDX_T, MemLayout::Scalar));
}
Op::Loop { len, .. } => {
*rcs.entry(len).or_insert(0) += 1;
dtypes.insert(op_id, (IDX_T, MemLayout::Scalar));
}
Op::If { condition } => {
*rcs.entry(condition).or_insert(0) += 1;
}
Op::Barrier | Op::EndIf | Op::EndLoop => {}
}
op_id = self.next_op(op_id);
}
if !op_id.is_null() {
panic!("compute_dtypes_and_rcs did not finish in 10000 steps");
}
(dtypes, rcs)
}
pub(crate) fn layout(&self, mut op_id: OpId) -> MemLayout {
for _ in 0..10000 {
match self.ops[op_id].op {
Op::Const(_) | Op::Param { .. } | Op::Storage { .. } => return MemLayout::Scalar,
Op::Range { .. } => return MemLayout::Scalar,
Op::Cast { x, .. } | Op::Bitcast { x, .. } => op_id = x,
Op::Load { layout, .. } => return layout,
Op::Store { src: x, .. } => op_id = x,
Op::Unary { x, .. } => op_id = x,
Op::Binary { x, .. } => op_id = x,
Op::Mad { x, .. } => op_id = x,
Op::Wmma { dims, .. } => match dims {
MMADims::m8n8k16 => return MemLayout::Vector(2),
MMADims::m16n8k8 => return MemLayout::Vector(4),
MMADims::m16n8k16 => return MemLayout::Vector(4),
MMADims::m32n8k16 => return MemLayout::Vector(8),
MMADims::m8n32k16 => return MemLayout::Vector(8),
MMADims::m8n8k32 => return MemLayout::Vector(2),
MMADims::m8n8k128 => return MemLayout::Vector(2),
},
Op::MatmulTile { acc, .. } => op_id = acc,
Op::TransposeTile { x } => op_id = x,
Op::Stack { ref ops } => {
return MemLayout::Vector(ops.len().try_into().unwrap());
}
Op::Asm { ref ops, .. } => op_id = ops[0],
Op::Index { .. } => return MemLayout::Scalar,
Op::Move { x, .. } => op_id = x,
Op::Reduce { x, .. } => op_id = x,
Op::ReduceTile { acc, .. } => op_id = acc,
Op::BroadcastTile { x, .. } => op_id = x,
Op::EndLoop | Op::Loop { .. } => return MemLayout::Scalar,
Op::Barrier | Op::If { .. } | Op::EndIf => todo!(),
}
}
panic!("layout not found for too long time");
}
pub(crate) fn dtype(&self, mut op_id: OpId) -> DType {
for _ in 0..10000 {
match self.ops[op_id].op {
Op::Const(c) => return c.dtype(),
Op::Param { dtype, .. } => return dtype,
Op::Storage { dtype, .. } => return dtype,
Op::Cast { dtype, .. } => return dtype,
Op::Bitcast { dtype, .. } => return dtype,
Op::Range { .. } => return IDX_T,
Op::Load { src, .. } => op_id = src,
Op::Unary { x, .. } => op_id = x,
Op::Binary { x, bop, .. } => {
if bop.returns_bool() {
return DType::Bool;
}
op_id = x;
}
Op::Mad { x, .. } => op_id = x,
Op::Wmma { dtype, .. } => match dtype {
MMADType::f16_f16_f16_f32 => return DType::F32,
MMADType::f16_f16_f16_f16 => return DType::F16,
MMADType::s8_s8_s32_s32
| MMADType::s4_s4_s32_s32
| MMADType::b1_b1_s32_xor_popc
| MMADType::b1_b1_s32_and_popc => return DType::I32,
},
Op::MatmulTile { acc, .. } => op_id = acc,
Op::TransposeTile { x } => op_id = x,
Op::Stack { ref ops } => op_id = ops[0],
Op::Asm { ref ops, .. } => op_id = ops[0],
Op::Index { vec, .. } => op_id = vec,
Op::Store { src: x, .. } => op_id = x,
Op::Move { x, .. } => op_id = x,
Op::Reduce { x, .. } => op_id = x,
Op::ReduceTile { acc, .. } => op_id = acc,
Op::BroadcastTile { x, .. } => op_id = x,
Op::EndLoop | Op::Loop { .. } => return IDX_T,
Op::Barrier | Op::If { .. } | Op::EndIf => todo!(),
}
}
panic!("dtype not found for too long time");
}
#[track_caller]
pub(crate) fn at(&self, op_id: OpId) -> &Op {
&self.ops[op_id].op
}
pub(crate) fn prev_op(&self, op_id: OpId) -> OpId {
self.ops[op_id].prev
}
pub(crate) fn next_op(&self, op_id: OpId) -> OpId {
self.ops[op_id].next
}
pub(crate) fn insert_before(&mut self, before_id: OpId, op: Op) -> OpId {
debug_assert!(!before_id.is_null());
debug_assert!(!self.ops.is_empty());
let prev = self.ops[before_id].prev;
let op_node = OpLinked { prev, next: before_id, op };
let op_id = self.ops.push(op_node);
self.ops[before_id].prev = op_id;
if prev.is_null() {
self.head = op_id;
} else {
self.ops[prev].next = op_id;
}
op_id
}
pub(crate) fn insert_const_idx_before(&mut self, before_id: OpId, val: impl crate::scalar::Scalar) -> OpId {
self.insert_before(before_id, Op::Const(Constant::idx(val)))
}
pub(crate) fn insert_after(&mut self, after_id: OpId, op: Op) -> OpId {
debug_assert!(!after_id.is_null());
debug_assert!(!self.ops.is_empty());
let next = self.ops[after_id].next;
let op_node = OpLinked { prev: after_id, next, op };
let op_id = self.ops.push(op_node);
self.ops[after_id].next = op_id;
if next.is_null() {
self.tail = op_id;
} else {
self.ops[next].prev = op_id;
}
op_id
}
pub(crate) fn move_op_after(&mut self, op_id: OpId, after_id: OpId) {
debug_assert!(!op_id.is_null());
debug_assert!(!after_id.is_null());
debug_assert!(!self.ops.is_empty());
if op_id == after_id {
return;
}
let OpLinked { prev, next, .. } = self.ops[op_id];
if prev.is_null() {
self.head = next;
} else {
self.ops[prev].next = next;
}
if next.is_null() {
self.tail = prev;
} else {
self.ops[next].prev = prev;
}
self.ops[op_id].prev = after_id;
let next = self.ops[after_id].next;
self.ops[op_id].next = next;
self.ops[after_id].next = op_id;
if next.is_null() {
self.tail = op_id;
} else {
self.ops[next].prev = op_id;
}
}
pub(crate) fn move_op_before(&mut self, op_id: OpId, before_id: OpId) {
debug_assert!(!op_id.is_null());
debug_assert!(!before_id.is_null());
debug_assert!(!self.ops.is_empty());
if op_id == before_id {
return;
}
let OpLinked { prev, next, .. } = self.ops[op_id];
if prev.is_null() {
self.head = next;
} else {
self.ops[prev].next = next;
}
if next.is_null() {
self.tail = prev;
} else {
self.ops[next].prev = prev;
}
self.ops[op_id].next = before_id;
let prev = self.ops[before_id].prev;
self.ops[op_id].prev = prev;
self.ops[before_id].prev = op_id;
if prev.is_null() {
self.head = op_id;
} else {
self.ops[prev].next = op_id;
}
}
pub(crate) fn remove_op(&mut self, op_id: OpId) {
debug_assert!(!op_id.is_null());
debug_assert!(!self.ops.is_empty());
let OpLinked { prev, next, .. } = self.ops[op_id];
if prev.is_null() {
self.head = next;
} else {
self.ops[prev].next = next;
}
if next.is_null() {
self.tail = prev;
} else {
self.ops[next].prev = prev;
}
self.ops.remove(op_id);
}
pub(crate) fn remove_unused_chain<T: Copy>(&mut self, x: OpId, keep_alive: &[OpId], loads: &[T]) -> Vec<T> {
let mut chain: Set<OpId> = Set::default();
let mut stack = vec![x];
for _ in 0..30_000 {
let Some(op) = stack.pop() else { break };
if chain.insert(op) {
stack.extend(self.ops[op].op.parameters().filter(|&p| !p.is_null()));
}
}
if !stack.is_empty() {
panic!("remove_unused_chain did not finish in 10000 steps");
}
let mut live: Set<OpId> = Set::default();
stack.extend_from_slice(keep_alive);
let mut op_id = self.head;
for _ in 0..30_000 {
if op_id.is_null() {
break;
}
if matches!(self.ops[op_id].op, Op::Store { .. }) {
stack.push(op_id);
}
op_id = self.next_op(op_id);
}
if !op_id.is_null() {
panic!("remove_unused_chain did not finish in 10000 steps");
}
for _ in 0..30_000 {
let Some(op) = stack.pop() else { break };
if live.insert(op) {
stack.extend(self.ops[op].op.parameters().filter(|&p| !p.is_null()));
}
}
if !stack.is_empty() {
panic!("remove_unused_chain did not finish in 10000 steps");
}
let param_ops: Vec<OpId> = {
let mut ops = Vec::new();
let mut id = self.head;
for _ in 0..30_000 {
if id.is_null() {
break;
}
if let Op::Param { kind, .. } = &self.ops[id].op {
if *kind != ParamKind::GlobalMut {
ops.push(id);
}
}
id = self.next_op(id);
}
if !id.is_null() {
panic!("remove_unused_chain did not finish in 10000 steps");
}
ops
};
let to_remove: Set<OpId> = chain.difference(&live).copied().collect();
let mut op_id = self.head;
for _ in 0..30_000 {
if op_id.is_null() {
break;
}
let next = self.next_op(op_id);
if to_remove.contains(&op_id) {
self.remove_op(op_id);
}
op_id = next;
}
if !op_id.is_null() {
panic!("remove_unused_chain did not finish in 10000 steps");
}
param_ops.iter().enumerate().filter(|&(_, &lv_id)| !to_remove.contains(&lv_id)).map(|(i, _)| loads[i]).collect()
}
pub(crate) fn iter_unordered(&self) -> impl Iterator<Item = (OpId, &Op)> {
self.ops.iter().map(|(id, node)| (id, &node.op))
}
pub(crate) fn name(&self) -> String {
let mut parts: Vec<&str> = Vec::new();
let mut op_id = self.head;
for _ in 0..10_000 {
if op_id.is_null() {
break;
}
match self.at(op_id) {
Op::Unary { uop, .. } => parts.push(match uop {
UOp::Neg => "neg",
UOp::Not => "not",
UOp::BitNot => "bitnot",
UOp::Exp => "exp",
UOp::Exp2 => "exp2",
UOp::Log2 => "log2",
UOp::Reciprocal => "reciprocal",
UOp::Sqrt => "sqrt",
UOp::Rsqrt => "rsqrt",
UOp::Sin => "sin",
UOp::Cos => "cos",
UOp::Floor => "floor",
UOp::Trunc => "trunc",
UOp::Abs => "abs",
}),
Op::Binary { bop, .. } => parts.push(match bop {
BOp::Add => "add",
BOp::Sub => "sub",
BOp::Mul => "mul",
BOp::Div => "div",
BOp::Pow => "pow",
BOp::Mod => "mod",
BOp::Cmplt => "cmplt",
BOp::Cmpgt => "cmpgt",
BOp::Cmpge => "cmpge",
BOp::Max => "max",
BOp::Or => "or",
BOp::And => "and",
BOp::BitXor => "bitxor",
BOp::BitOr => "bitor",
BOp::BitAnd => "bitand",
BOp::BitShiftLeft => "shl",
BOp::BitShiftRight => "shr",
BOp::NotEq => "neq",
BOp::Eq => "eq",
}),
Op::Reduce { rop, .. } => parts.push(match rop {
BOp::Add => "sum",
BOp::Max => "max",
BOp::Mul => "prod",
_ => "reduce",
}),
Op::ReduceTile { rop, .. } => parts.push(match rop {
BOp::Add => "reduce_tile_sum",
BOp::Max => "reduce_tile_max",
BOp::Mul => "reduce_tile_prod",
_ => "reduce_tile",
}),
Op::Mad { .. } => parts.push("mad"),
Op::Wmma { .. } => parts.push("wmma"),
Op::Cast { .. } => parts.push("cast"),
Op::Bitcast { .. } => parts.push("bitcast"),
_ => {}
}
op_id = self.next_op(op_id);
}
if !op_id.is_null() {
panic!("name did not finish in 10000 steps");
}
parts.dedup();
if parts.is_empty() {
return "copy".into();
}
parts.join("_")
}
pub(crate) fn contains_stores(&self) -> bool {
self.ops.values().any(|x| matches!(x.op, Op::Store { .. }))
}
pub fn flop_mem_rw(&self) -> (u64, u64, u64) {
let mut flops: u64 = 0;
let mut read: u64 = 0;
let mut write: u64 = 0;
let mut loop_stack: Vec<u64> = Vec::new();
let mut group_stack: Vec<u64> = Vec::new();
let mut local_prod: u64 = 1;
let mut op_id = self.head;
for _ in 0..20_000 {
if op_id.is_null() {
break;
}
let loop_prod: u64 = loop_stack.iter().product::<u64>().max(1);
let group_prod: u64 = group_stack.iter().product::<u64>().max(1);
let mult: u64 = loop_prod.saturating_mul(group_prod).saturating_mul(local_prod);
match self.at(op_id) {
Op::Loop { len } => {
let v = self.resolve_const(*len).and_then(|c| c.as_dim()).unwrap_or(1).max(1) as u64;
loop_stack.push(v);
}
Op::EndLoop => {
loop_stack.pop();
}
Op::Range { kind, .. } => match kind {
RangeKind::Group(l) => {
let v = self.resolve_const(*l).and_then(|c| c.as_dim()).unwrap_or(1).max(1) as u64;
group_stack.push(v);
}
RangeKind::Local(n) => {
local_prod = local_prod.saturating_mul(*n as u64);
}
RangeKind::Warp(_) => {}
},
Op::Unary { .. } => {
if self.dtype(op_id).is_float() {
flops = flops.saturating_add(mult)
}
}
Op::Binary { .. } => {
if self.dtype(op_id).is_float() {
flops = flops.saturating_add(mult)
}
}
Op::Mad { .. } => {
if self.dtype(op_id).is_float() {
flops = flops.saturating_add(2 * mult)
}
}
Op::Wmma { dims, .. } => {
let (m, n, k) = dims.decompose_mnk();
let ws = u64::from(self.dev_info().warp_size);
let warps = (group_prod * local_prod / ws.max(1)).max(1);
let loop_prod: u64 = loop_stack.iter().product::<u64>().max(1);
flops = flops.saturating_add(2 * m * n * k * warps * loop_prod);
}
Op::ReduceTile { .. } => flops = flops.saturating_add(mult),
Op::Load { src, layout, .. } => {
if let Op::Param { kind: ParamKind::Global, dtype, .. } = &self.ops[*src].op {
let bytes = (dtype.bit_size() as u64 / 8) * layout.n_elements() as u64;
read = read.saturating_add(bytes.saturating_mul(mult));
}
}
Op::Store { dst, layout, .. } => {
if let Op::Param { kind: ParamKind::GlobalMut, dtype, .. } = &self.ops[*dst].op {
let bytes = (dtype.bit_size() as u64 / 8) * layout.n_elements() as u64;
write = write.saturating_add(bytes.saturating_mul(mult));
}
}
_ => {}
}
op_id = self.next_op(op_id);
}
(flops, read, write)
}
pub(crate) fn is_reduce(&self) -> bool {
self.ops.values().any(|x| matches!(x.op, Op::Reduce { .. } | Op::ReduceTile { .. }))
}
#[must_use]
pub(crate) fn shape_ids(&mut self, op_id: OpId) -> Vec<OpId> {
if op_id.is_null() {
return Vec::new();
}
if let Some(cached) = self.shape_cache.get(&op_id) {
return cached.clone();
}
fn descriptor(k: &mut Kernel, id: OpId) -> Vec<OpId> {
if id.is_null() {
return Vec::new();
}
let (shape, dtype) = match k.ops[id].op {
Op::Stack { ref ops } => return ops.to_vec(),
Op::Const(_) => return vec![id],
Op::Unary { .. } | Op::Binary { .. } | Op::Load { .. } => return vec![id],
Op::Param { shape, dtype, .. } => (shape, dtype),
ref op => todo!("shape_ids: invalid shape descriptor {op:?}"),
};
debug_assert!(dtype == IDX_T, "shape tensor must be {IDX_T:?}, got {dtype:?}");
if shape.is_null() {
return vec![id];
}
let rank = k.shape(shape).len();
debug_assert!(rank <= 1, "shape_ids: param shape descriptor must be 0d or 1d, got rank {rank}");
let mut dims = Vec::with_capacity(rank);
for i in 0..rank {
dims.push(k.push_back(Op::Index { vec: id, idx: i }));
}
dims
}
let root = op_id;
let mut stack = vec![op_id];
let mut visited: Map<OpId, Vec<OpId>> = Map::default();
for _ in 0..10_000 {
let Some(op_id) = stack.pop() else {
let shape = visited.remove(&op_id).unwrap();
self.shape_cache.insert(root, shape.clone());
return shape;
};
if visited.contains_key(&op_id) {
continue;
}
let node_op = self.ops[op_id].op.clone();
match node_op {
Op::Const(_) => {
visited.insert(op_id, vec![]);
}
Op::Param { shape, .. } => {
visited.insert(op_id, descriptor(self, shape));
}
Op::Stack { ref ops } => {
if ops.is_empty() {
visited.insert(op_id, vec![]);
} else if let Some(mut dims) = visited.get(&ops[0]).cloned() {
dims.insert(0, self.const_idx(ops.len() as u32));
visited.insert(op_id, dims);
} else {
stack.push(op_id);
stack.push(ops[0]);
}
}
Op::Move { x, ref mop } => match mop.as_ref() {
MoveOp::Reshape { shape, .. } | MoveOp::Expand { shape } => {
visited.insert(op_id, descriptor(self, *shape));
}
MoveOp::Permute { axes } => match visited.get(&x) {
Some(dims) => {
visited.insert(op_id, crate::shape::permute(dims, axes));
}
None => {
stack.push(op_id);
stack.push(x);
}
},
MoveOp::Flip { .. } => match visited.get(&x) {
Some(dims) => {
visited.insert(op_id, dims.clone());
}
None => {
stack.push(op_id);
stack.push(x);
}
},
&MoveOp::Narrow { axis, len, .. } | &MoveOp::Pad { axis, len, .. } => match visited.get(&x).cloned() {
Some(mut dims) => {
if dims.is_empty() {
dims = vec![len];
} else {
dims[axis as usize] = len;
}
visited.insert(op_id, dims);
}
None => {
stack.push(op_id);
stack.push(x);
}
},
},
Op::Binary { x, y, .. } => {
let x_done = visited.contains_key(&x);
if x_done && !visited[&x].is_empty() {
visited.insert(op_id, visited[&x].clone());
} else if visited.contains_key(&y) {
visited.insert(op_id, visited[&y].clone());
} else {
stack.push(op_id);
if !x_done {
stack.push(x);
}
stack.push(y);
}
}
Op::Reduce { x, .. } => match visited.get(&x).cloned() {
Some(mut dims) => {
dims.truncate(dims.len().saturating_sub(1));
visited.insert(op_id, dims);
}
None => {
stack.push(op_id);
stack.push(x);
}
},
Op::Load { src: x, .. }
| Op::Store { src: x, .. }
| Op::Cast { x, .. }
| Op::Bitcast { x, .. }
| Op::Unary { x, .. }
| Op::Mad { x, .. }
| Op::MatmulTile { x, .. }
| Op::ReduceTile { x, .. }
| Op::Index { vec: x, .. }
| Op::TransposeTile { x } => match visited.get(&x) {
Some(dims) => {
visited.insert(op_id, dims.clone());
}
None => {
stack.push(op_id);
stack.push(x);
}
},
Op::Asm { ref ops, .. } => match visited.get(&ops[0]) {
Some(dims) => {
visited.insert(op_id, dims.clone());
}
None => {
stack.push(op_id);
stack.push(ops[0]);
}
},
Op::Range { .. } | Op::Loop { .. } => {
visited.insert(op_id, vec![]);
}
Op::Storage { len, .. } if len == 1 => {
visited.insert(op_id, vec![]);
}
ref op => todo!("shape_ids of {op:?}"),
}
}
panic!("shape_ids did not resolve in 10000 steps");
}
pub(crate) fn shape(&self, op_id: OpId) -> Vec<Dim> {
fn descriptor(k: &Kernel, id: OpId) -> Vec<OpId> {
if id.is_null() {
return Vec::new();
}
match k.ops[id].op {
Op::Stack { ref ops } => ops.to_vec(),
_ => vec![id],
}
}
if op_id.is_null() {
return Vec::new();
}
let mut stack = vec![op_id];
let mut visited: Map<OpId, Vec<Dim>> = Map::default();
for _ in 0..10_000 {
let Some(id) = stack.pop() else {
return visited[&op_id].clone();
};
if visited.contains_key(&id) {
continue;
}
match self.ops[id].op {
Op::Const(_) => {
visited.insert(id, vec![]);
}
Op::Param { shape, .. } => {
visited.insert(
id,
descriptor(self, shape)
.iter()
.map(|&d| self.resolve_const(d).and_then(crate::dtype::Constant::as_dim).unwrap_or(-1))
.collect(),
);
}
Op::Stack { ref ops } => {
if ops.is_empty() {
visited.insert(id, vec![]);
} else if let Some(mut dims) = visited.get(&ops[0]).cloned() {
dims.insert(0, ops.len() as Dim);
visited.insert(id, dims);
} else {
stack.push(id);
stack.push(ops[0]);
}
}
Op::Move { x, ref mop } => match mop.as_ref() {
MoveOp::Reshape { shape, .. } | MoveOp::Expand { shape } => {
visited.insert(
id,
descriptor(self, *shape)
.iter()
.map(|&d| self.resolve_const(d).and_then(crate::dtype::Constant::as_dim).unwrap_or(-1))
.collect(),
);
}
MoveOp::Permute { axes } => match visited.get(&x) {
Some(dims) => {
visited.insert(id, crate::shape::permute(dims, axes));
}
None => {
stack.push(id);
stack.push(x);
}
},
MoveOp::Flip { .. } => match visited.get(&x) {
Some(dims) => {
visited.insert(id, dims.clone());
}
None => {
stack.push(id);
stack.push(x);
}
},
&MoveOp::Narrow { axis, len, .. } | &MoveOp::Pad { axis, len, .. } => match visited.get(&x).cloned() {
Some(mut dims) => {
if dims.is_empty() {
dims = vec![self.resolve_const(len).and_then(crate::dtype::Constant::as_dim).unwrap_or(-1)];
} else {
dims[axis as usize] =
self.resolve_const(len).and_then(crate::dtype::Constant::as_dim).unwrap_or(-1);
}
visited.insert(id, dims);
}
None => {
stack.push(id);
stack.push(x);
}
},
},
Op::Binary { x, y, .. } => {
let x_done = visited.contains_key(&x);
if x_done && !visited[&x].is_empty() {
visited.insert(id, visited[&x].clone());
} else if visited.contains_key(&y) {
visited.insert(id, visited[&y].clone());
} else {
stack.push(id);
if !x_done {
stack.push(x);
}
stack.push(y);
}
}
Op::Reduce { x, .. } => match visited.get(&x).cloned() {
Some(mut dims) => {
dims.truncate(dims.len().saturating_sub(1));
visited.insert(id, dims);
}
None => {
stack.push(id);
stack.push(x);
}
},
Op::Load { src: x, .. }
| Op::Store { src: x, .. }
| Op::Cast { x, .. }
| Op::Bitcast { x, .. }
| Op::Unary { x, .. }
| Op::Mad { x, .. }
| Op::MatmulTile { x, .. }
| Op::ReduceTile { x, .. }
| Op::Index { vec: x, .. }
| Op::TransposeTile { x } => match visited.get(&x) {
Some(dims) => {
visited.insert(id, dims.clone());
}
None => {
stack.push(id);
stack.push(x);
}
},
Op::Asm { ref ops, .. } => match visited.get(&ops[0]) {
Some(dims) => {
visited.insert(id, dims.clone());
}
None => {
stack.push(id);
stack.push(ops[0]);
}
},
Op::Range { .. } | Op::Loop { .. } => {
visited.insert(id, vec![]);
}
Op::Storage { len, .. } if len == 1 => {
visited.insert(id, vec![]);
}
ref op => todo!("shape of {op:?}"),
}
}
panic!("shape did not resolve in 10000 steps");
}
#[must_use]
pub(crate) fn stack_shape_dims(&mut self, op_id: OpId) -> OpId {
match self.shape_ids(op_id).as_slice() {
[] => OpId::NULL,
[dim] => *dim,
dims => self.stack(dims),
}
}
pub(crate) fn get_strides(&self, index: OpId) -> Map<OpId, (Dim, Dim)> {
let index_len_of = |op: &Op| -> Dim {
match op {
Op::Range { kind, .. } => match kind {
RangeKind::Group(len) => self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap_or(42),
RangeKind::Local(len) => i64::from(*len),
RangeKind::Warp(_) => i64::from(self.dev_info().warp_size),
},
_ => unreachable!(),
}
};
let mut params = vec![(index, 1i64)];
let mut indices = Map::default();
for _ in 0..10_000 {
let Some((param, scale)) = params.pop() else { break };
match self.ops[param].op {
Op::Binary { x, y, bop } => {
if bop == BOp::Add {
if let Op::Loop { len, .. } = self.ops[x].op {
indices.insert(x, (self.resolve_const(len).and_then(crate::dtype::Constant::as_dim).unwrap(), 1));
params.push((y, scale));
} else if let Op::Range { .. } = self.ops[x].op {
indices.insert(x, (index_len_of(&self.ops[x].op), 1));
params.push((y, scale));
} else if let Op::Loop { len, .. } = self.ops[y].op {
indices.insert(y, (self.resolve_const(len).and_then(crate::dtype::Constant::as_dim).unwrap(), 1));
params.push((x, scale));
} else if let Op::Range { .. } = self.ops[y].op {
indices.insert(y, (index_len_of(&self.ops[y].op), 1));
params.push((x, scale));
} else {
params.push((x, scale));
params.push((y, scale));
}
}
if bop == BOp::Mul {
match (&self.ops[x].op, &self.ops[y].op) {
(Op::Loop { len, .. }, Op::Const(c)) => {
indices.insert(
x,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
c.as_dim().unwrap() * scale,
),
);
}
(Op::Const(c), Op::Loop { len, .. }) => {
indices.insert(
y,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
c.as_dim().unwrap() * scale,
),
);
}
(Op::Range { .. }, Op::Const(c)) => {
indices.insert(x, (index_len_of(&self.ops[x].op), c.as_dim().unwrap() * scale));
}
(Op::Const(c), Op::Range { .. }) => {
indices.insert(y, (index_len_of(&self.ops[y].op), c.as_dim().unwrap() * scale));
}
_ => {}
}
}
if bop == BOp::BitShiftLeft {
match (&self.ops[x].op, &self.ops[y].op) {
(Op::Loop { len, .. }, Op::Const(c)) => {
indices.insert(
x,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
(1i64 << c.as_dim().unwrap()) * scale,
),
);
}
(Op::Range { .. }, Op::Const(c)) => {
indices.insert(x, (index_len_of(&self.ops[x].op), (1i64 << c.as_dim().unwrap()) * scale));
}
(Op::Const(c), Op::Range { .. }) => {
indices.insert(y, (index_len_of(&self.ops[y].op), (1i64 << c.as_dim().unwrap()) * scale));
}
(Op::Const(c), Op::Loop { len, .. }) => {
indices.insert(
y,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
(1i64 << c.as_dim().unwrap()) * scale,
),
);
}
_ => {
if let Op::Const(c) = self.ops[y].op {
params.push((x, scale * (1i64 << c.as_dim().unwrap())));
}
}
}
}
}
Op::Mad { x, y, z } => {
match &self.ops[z].op {
Op::Loop { len, .. } => {
indices.insert(z, (self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(), 1));
}
Op::Range { .. } => {
indices.insert(z, (index_len_of(&self.ops[z].op), 1));
}
_ => {
params.push((z, scale));
}
}
match (&self.ops[x].op, &self.ops[y].op) {
(Op::Loop { len, .. }, Op::Const(c)) => {
indices.insert(
x,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
c.as_dim().unwrap() * scale,
),
);
}
(Op::Range { .. }, Op::Const(c)) => {
indices.insert(x, (index_len_of(&self.ops[x].op), c.as_dim().unwrap() * scale));
}
(Op::Const(c), Op::Loop { len, .. }) => {
indices.insert(
y,
(
self.resolve_const(*len).and_then(crate::dtype::Constant::as_dim).unwrap(),
c.as_dim().unwrap() * scale,
),
);
}
(Op::Const(c), Op::Range { .. }) => {
indices.insert(y, (index_len_of(&self.ops[y].op), c.as_dim().unwrap() * scale));
}
_ => {}
}
}
Op::Const(c) => {
indices
.entry(OpId::NULL)
.and_modify(|(_, v)| *v += c.as_dim().unwrap() * scale)
.or_insert((0, c.as_dim().unwrap() * scale));
}
_ => {}
}
}
if !params.is_empty() {
panic!("get_strides did not finish in 10000 steps");
}
indices
}
pub(crate) fn resolve_const(&self, op_id: OpId) -> Option<Constant> {
let mut seen: Set<OpId> = Set::default();
let mut order: Vec<OpId> = Vec::new();
let mut stack = vec![(op_id, false)];
for _ in 0..10_000 {
let Some((id, emit)) = stack.pop() else { break };
if id.is_null() {
continue;
}
if emit {
order.push(id);
continue;
}
if !seen.insert(id) {
continue;
}
stack.push((id, true));
match &self.ops[id].op {
Op::Cast { x, .. } => stack.push((*x, false)),
Op::Bitcast { x, .. } => stack.push((*x, false)),
Op::Unary { x, .. } => stack.push((*x, false)),
Op::Binary { x, y, .. } => {
stack.push((*x, false));
stack.push((*y, false));
}
Op::Stack { ops } => stack.extend(ops.iter().copied().map(|o| (o, false))),
Op::Loop { len } => stack.push((*len, false)),
&Op::Range { kind, .. } => match kind {
RangeKind::Group(len) => stack.push((len, false)),
RangeKind::Local(_) | RangeKind::Warp(_) => {}
},
Op::Mad { x, y, z } => {
stack.push((*x, false));
stack.push((*y, false));
stack.push((*z, false));
}
_ => {}
}
}
if !stack.is_empty() {
panic!("resolve_const did not finish in 10000 steps");
}
let mut values: Map<OpId, Option<Constant>> = Map::default();
for &id in &order {
let v = match &self.ops[id].op {
Op::Const(c) => Some(*c),
Op::Cast { x, dtype } => values.get(x).copied().flatten().map(|v| v.cast(*dtype)),
Op::Bitcast { x, dtype } => values.get(x).copied().flatten().map(|v| v.bitcast(*dtype)),
Op::Unary { x, uop } => values.get(x).copied().flatten().map(|v| v.unary(*uop)),
Op::Binary { x, y, bop } => {
values.get(x).copied().flatten().zip(values.get(y).copied().flatten()).map(|(a, b)| {
let dt = a.dtype().least_upper_dtype(b.dtype());
Constant::binary(a.cast(dt), b.cast(dt), *bop)
})
}
Op::Mad { x, y, z } => values
.get(x)
.copied()
.flatten()
.zip(values.get(y).copied().flatten())
.zip(values.get(z).copied().flatten())
.map(|((a, b), c)| {
let dt = a.dtype().least_upper_dtype(b.dtype()).least_upper_dtype(c.dtype());
let prod = Constant::binary(a.cast(dt), b.cast(dt), BOp::Mul);
Constant::binary(prod, c.cast(dt), BOp::Add)
}),
Op::Loop { len } => values.get(len).copied().flatten(),
&Op::Range { kind, .. } => match kind {
RangeKind::Group(len) => values.get(&len).copied().flatten(),
RangeKind::Local(len) => Some(Constant::idx(len)),
RangeKind::Warp(_) => None,
},
Op::Param { .. } => None,
Op::Load { .. } => None,
ref op => todo!("{op:?}"),
};
values.insert(id, v);
}
values[&op_id]
}
fn remap(&mut self, x: OpId, y: OpId) {
for op_node in self.ops.values_mut() {
for param in op_node.op.parameters_mut() {
if *param == x {
*param = y;
}
}
}
}
pub(crate) fn push_back(&mut self, op: Op) -> OpId {
let op_node = OpLinked { prev: self.tail, next: OpId::NULL, op };
let op_id = self.ops.push(op_node);
if self.head.is_null() {
self.head = op_id;
} else {
self.ops[self.tail].next = op_id;
}
self.tail = op_id;
op_id
}
pub(crate) fn duplicate_subkernel<T: Copy>(&mut self, root_op: OpId, loads: &[T]) -> (Self, OpId, Vec<T>) {
let mut root_required = Set::default();
let mut stack = vec![root_op];
for _ in 0..10_000 {
let Some(op) = stack.pop() else { break };
if root_required.insert(op) {
stack.extend(self.at(op).parameters());
}
}
if !stack.is_empty() {
panic!("duplicate_subkernel did not finish in 10000 steps");
}
let mut new_loads: Vec<T> = Vec::new();
let mut load_idx = 0;
let mut oid = self.head;
for _ in 0..10_000 {
if oid.is_null() {
break;
}
match self.at(oid) {
Op::Param { kind: ParamKind::Global | ParamKind::Variable, .. } => {
if root_required.contains(&oid) {
new_loads.push(loads[load_idx]);
}
load_idx += 1;
}
Op::Param { kind: ParamKind::GlobalMut, .. } => {}
Op::Storage { .. } => {
panic!("duplicate_subkernel: unexpected Op::Storage (pre-linearize kernels contain only Params)")
}
_ => {}
}
oid = self.next_op(oid);
}
if !oid.is_null() {
panic!("duplicate_subkernel did not finish in 10000 steps");
}
let mut new_kernel = Kernel::from_device_id(self.dev, self.dev_info.clone());
let mut remap: Map<OpId, OpId> =
Map::with_capacity_and_hasher(root_required.len(), core::hash::BuildHasherDefault::default());
let mut new_root_op = OpId::NULL;
let mut old_id = self.head;
for _ in 0..10_000 {
if old_id.is_null() {
break;
}
if root_required.contains(&old_id) {
let mut op = self.at(old_id).clone();
op.remap_params(&remap);
let new_id = new_kernel.push_back(op);
if old_id == root_op {
new_root_op = new_id;
}
remap.insert(old_id, new_id);
}
old_id = self.next_op(old_id);
}
if !old_id.is_null() {
panic!("duplicate_subkernel did not finish in 10000 steps");
}
(new_kernel, new_root_op, new_loads)
}
pub(crate) fn get_group_indices(&self) -> std::collections::BTreeMap<u32, OpId> {
let mut indices = std::collections::BTreeMap::new();
for (op_id, op_node) in self.ops.iter() {
if let Op::Range { axis, kind: RangeKind::Group(_), .. } = op_node.op {
indices.insert(axis, op_id);
}
}
indices
}
pub(crate) fn renumber_indices(&mut self) {
let mut group_indices = BTreeMap::default();
let mut local_indices = BTreeMap::default();
for (op_id, op_node) in self.ops.iter() {
match op_node.op {
Op::Range { axis, kind: RangeKind::Group(_), .. } => group_indices.insert(axis, op_id),
Op::Range { axis, kind: RangeKind::Local(_), .. } => local_indices.insert(axis, op_id),
_ => None,
};
}
let mut ax = 0;
for &idx_id in group_indices.values() {
let Op::Range { axis, kind: RangeKind::Group(_), .. } = &mut self.ops[idx_id].op else {
unreachable!()
};
*axis = ax;
ax += 1;
}
for &idx_id in local_indices.values() {
let Op::Range { axis, kind: RangeKind::Local(_), .. } = &mut self.ops[idx_id].op else {
unreachable!()
};
*axis = ax;
ax += 1;
}
}
pub(crate) fn is_preceded_by_reduce(&self, x: OpId) -> bool {
let mut params = vec![x];
let mut found = false;
let mut seen: Set<OpId> = Set::default();
for _ in 0..10_000 {
let Some(param) = params.pop() else { break };
if !seen.insert(param) {
continue;
}
if let &Op::Reduce { x, .. } = self.at(param) {
params = vec![x];
found = true;
break;
}
params.extend(self.ops[param].op.parameters());
}
if !found && !params.is_empty() {
panic!("is_preceded_by_reduce did not finish in 10000 steps");
}
let mut seen: Set<OpId> = Set::default();
for _ in 0..10_000 {
let Some(param) = params.pop() else { break };
if !seen.insert(param) {
continue;
}
if matches!(self.ops[param].op, Op::Param { .. } | Op::Reduce { .. }) {
return true;
}
params.extend(self.ops[param].op.parameters());
}
if !params.is_empty() {
panic!("is_preceded_by_reduce did not finish in 10000 steps");
}
false
}
pub(crate) fn is_preceded_by_compute(&self, x: OpId) -> bool {
let mut params = vec![x];
let mut seen: Set<OpId> = Set::default();
let (mut has_compute, mut has_param) = (false, false);
for _ in 0..10_000 {
let Some(param) = params.pop() else { break };
if !seen.insert(param) {
continue;
}
match &self.ops[param].op {
Op::Binary { .. } | Op::Unary { .. } | Op::Reduce { .. } => {
has_compute = true;
params.extend(self.ops[param].op.parameters());
}
Op::Param { .. } => has_param = true,
Op::Const(_) => {}
_ => params.extend(self.ops[param].op.parameters()),
}
}
if !params.is_empty() {
panic!("is_preceded_by_compute did not finish in 10000 steps");
}
has_compute && has_param
}
}