use std::{collections::BTreeSet, hash::BuildHasherDefault, path::Path};
#[cfg(feature = "viz")]
use crate::viz::Viz;
use crate::{
DType, Dev, Map, Scalar, Set, ZyxError,
backend::{Buffer, DTypeCapability, DeviceProgramId, LaunchArg, Pool, ProgramId},
dtype::Constant,
graph::{ExecPlan, Graph, GraphId, Node},
kernel::{BOp, Kernel, MoveOp, Op, OpId, ParamKind, UOp},
rng::Rng,
scalar::{bf16, f8e4m3, f8e5m2, f16},
shape::{Dim, UAxis},
slab::{Slab, SlabId},
symbolic::{Expr, ExprId},
tensor::TensorId,
};
pub fn loads_dropped_by_prune(old: &[TensorId], new: &[TensorId]) -> Vec<TensorId> {
let mut dropped = Vec::new();
let mut seen: Set<TensorId> = Set::default();
for &tid in old {
if !seen.insert(tid) {
continue;
}
let old_c = old.iter().filter(|&&t| t == tid).count();
let new_c = new.iter().filter(|&&t| t == tid).count();
dropped.extend(std::iter::repeat_n(tid, old_c - new_c));
}
dropped
}
#[derive(Debug, Clone, Copy, PartialEq, PartialOrd, Eq, Ord, Hash)]
pub struct KernelId(u16);
impl From<usize> for KernelId {
fn from(value: usize) -> Self {
KernelId(value as u16)
}
}
impl From<KernelId> for usize {
fn from(value: KernelId) -> Self {
value.0 as usize
}
}
impl SlabId for KernelId {
const ZERO: Self = Self(0);
const NULL: Self = Self(u16::MAX);
fn inc(&mut self) {
self.0 += 1;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum ResolvedDim {
Static(Dim),
Symbolic(ExprId),
}
#[derive(Debug)]
pub enum TensorData {
Leaf {
shape_id: ExprId,
dtype: DType,
buffer: Buffer,
rc: u16,
},
PendingLeaf {
old_buffer: Option<Buffer>,
depends_on: KernelId,
shape_id: ExprId,
dtype: DType,
rc: u16,
},
GraphLeaf {
class_id: OpId,
graph_id: GraphId,
shape_id: ExprId,
dtype: DType,
rc: u16,
buffer: Buffer,
},
Eager {
kernel_id: KernelId,
op_id: OpId,
shape_id: ExprId,
dtype: DType,
rc: u16,
},
Graph {
class_id: OpId,
graph_id: GraphId,
shape_id: ExprId,
dtype: DType,
rc: u16,
},
Promoted {
kernel_id: KernelId,
op_id: OpId,
class_id: OpId,
graph_id: GraphId,
shape_id: ExprId,
dtype: DType,
rc: u16,
},
Symbolic {
expr: ExprId,
rc: u16,
},
}
#[derive(Debug)]
pub struct KernelData {
pub outputs: Set<TensorId>,
pub loads: Vec<TensorId>,
pub stores: Vec<TensorId>,
pub kernel: Kernel,
}
pub struct Runtime {
pub graphs: Slab<GraphId, Graph>,
pub tensors: Slab<TensorId, TensorData>,
pub kernels: Slab<KernelId, KernelData>,
kernel_map: Map<Kernel, KernelId>,
programs: Map<KernelId, DeviceProgramId>,
timings: Map<ProgramId, u64>,
pub(crate) expr_hash: Map<Expr, ExprId>,
pub exprs: Slab<ExprId, Expr>,
pub rng: Rng,
pub implicit_casts: bool,
pub training: bool,
pub plan_cache: Map<u64, ExecPlan>,
#[cfg(feature = "viz")]
pub viz: Viz,
}
impl Runtime {
pub(crate) fn plan_cache_key(&self, graph_id: GraphId, outputs: &BTreeSet<OpId>) -> u64 {
use std::hash::{Hash, Hasher};
let graph = &self.graphs[graph_id];
let mut hasher = std::collections::hash_map::DefaultHasher::new();
graph.cache_key(outputs).hash(&mut hasher);
for &cid in &graph.leaf_classes {
let &tid = graph.leaf_map.get(&cid).unwrap();
self.leaf_buffer(tid).map(|b| b.pool).hash(&mut hasher);
}
hasher.finish()
}
pub const fn new() -> Self {
Runtime {
graphs: Slab::new(),
tensors: Slab::new(),
kernels: Slab::new(),
kernel_map: Map::with_hasher(BuildHasherDefault::new()),
programs: Map::with_hasher(BuildHasherDefault::new()),
timings: Map::with_hasher(BuildHasherDefault::new()),
expr_hash: Map::with_hasher(BuildHasherDefault::new()),
exprs: Slab::new(),
rng: Rng::seed_from_u64(42069),
implicit_casts: true,
training: false,
plan_cache: Map::with_hasher(BuildHasherDefault::new()),
#[cfg(feature = "viz")]
viz: Viz::new(),
}
}
pub fn dtype(&self, x: TensorId) -> DType {
match self.tensors[x] {
TensorData::Eager { dtype, .. }
| TensorData::Promoted { dtype, .. }
| TensorData::Graph { dtype, .. }
| TensorData::PendingLeaf { dtype, .. }
| TensorData::Leaf { dtype, .. }
| TensorData::GraphLeaf { dtype, .. } => dtype,
TensorData::Symbolic { expr, .. } => self.expr_dtype(expr),
}
}
pub fn is_realized(&self, x: TensorId) -> bool {
self.leaf_buffer(x).is_some() || self.resolve_symbolic(x).is_some()
}
pub(crate) fn is_graph(&self, x: TensorId) -> bool {
match self.tensors[x] {
TensorData::GraphLeaf { .. } | TensorData::Graph { .. } | TensorData::Promoted { .. } => true,
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } | TensorData::Symbolic { .. } => {
false
}
}
}
pub fn supports_dtype(&mut self, dtype: DType) -> DTypeCapability {
let mut caps = DTypeCapability::none();
for dev in Dev::all() {
caps = caps.include(dev.info().supports_dtype(dtype));
}
caps
}
pub fn retain(&mut self, x: TensorId) {
if !x.is_null() {
match &mut self.tensors[x] {
TensorData::Eager { rc, .. }
| TensorData::PendingLeaf { rc, .. }
| TensorData::Leaf { rc, .. }
| TensorData::GraphLeaf { rc, .. }
| TensorData::Graph { rc, .. }
| TensorData::Promoted { rc, .. }
| TensorData::Symbolic { rc, .. } => {
*rc += 1;
#[cfg(feature = "debug_tensor_op")]
println!("rc::retain({x}) -> {rc}");
}
};
}
}
pub(crate) fn leaf_buffer(&self, x: TensorId) -> Option<Buffer> {
match self.tensors[x] {
TensorData::Leaf { buffer, .. } | TensorData::GraphLeaf { buffer, .. } => Some(buffer),
TensorData::PendingLeaf { .. }
| TensorData::Eager { .. }
| TensorData::Graph { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => None,
}
}
pub fn release(&mut self, x: TensorId) {
#[cfg(feature = "debug_tensor_op")]
{
let desc: String = match &self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => format!("eager kernel={kernel_id:?} op={op_id:?}"),
TensorData::PendingLeaf { shape_id, .. } => format!("pending shape={shape_id:?}"),
TensorData::Leaf { shape_id, buffer, .. } => {
format!("leaf shape={shape_id:?} buffer={buffer:?}")
}
TensorData::GraphLeaf { shape_id, buffer, .. } => {
format!("graphleaf shape={shape_id:?} buffer={buffer:?}")
}
TensorData::Graph { class_id, graph_id, .. } => format!("graph class={class_id:?} graph={graph_id:?}"),
TensorData::Promoted { kernel_id, class_id, graph_id, .. } => {
format!("promoted kernel={kernel_id:?} class={class_id:?} graph={graph_id:?}")
}
TensorData::Symbolic { expr, .. } => format!("symbolic {expr:?}"),
};
println!("runtime::release(tid={x}) kind={desc}");
}
let rc = {
match &mut self.tensors[x] {
TensorData::Promoted { rc, .. }
| TensorData::Eager { rc, .. }
| TensorData::Leaf { rc, .. }
| TensorData::PendingLeaf { rc, .. }
| TensorData::GraphLeaf { rc, .. }
| TensorData::Graph { rc, .. }
| TensorData::Symbolic { rc, .. } => {
*rc -= 1;
*rc
}
}
};
#[cfg(feature = "debug_tensor_op")]
println!("rc::release({x}) -> rc={rc}");
if rc != 0 {
return;
}
match self.tensors[x] {
TensorData::Symbolic { .. } => {
self.tensors.remove(x);
}
TensorData::GraphLeaf { buffer: buffer_id, .. } => {
buffer_id.pool.release(buffer_id.buffer_id);
self.tensors.remove(x);
}
TensorData::Leaf { buffer: buffer_id, .. } => {
buffer_id.pool.release(buffer_id.buffer_id);
self.tensors.remove(x);
}
TensorData::PendingLeaf { depends_on, old_buffer, .. } => {
debug_assert!(
old_buffer.is_none(),
"release: PendingLeaf {x} dies with a kept buffer — its store kernel never launched"
);
if !depends_on.is_null() && self.kernels.contains_id(depends_on) {
let mut_params: Vec<OpId> = {
let kd = &self.kernels[depends_on];
let mut mut_params: Vec<OpId> = Vec::new();
let mut i = kd.kernel.head;
for _ in 0..100_000 {
if i.is_null() {
break;
}
if matches!(kd.kernel.ops[i].op, Op::Param { kind: ParamKind::GlobalMut, .. }) {
mut_params.push(i);
}
i = kd.kernel.next_op(i);
}
debug_assert_eq!(
mut_params.len(),
kd.stores.len(),
"GlobalMut params and stores vec diverged in {depends_on:?}"
);
mut_params
};
let dead_params: Vec<OpId> = {
let kd = &self.kernels[depends_on];
mut_params.iter().enumerate().filter(|(idx, _)| kd.stores[*idx] == x).map(|(_, op)| *op).collect()
};
if !dead_params.is_empty() {
let mut keep_alive: Vec<OpId> = {
let kd = &self.kernels[depends_on];
mut_params
.iter()
.enumerate()
.filter(|&(idx, _)| kd.stores[idx] != x)
.flat_map(|(_, param)| {
let mut stores_to_param: Vec<OpId> = Vec::new();
let mut i = kd.kernel.head;
for _ in 0..100_000 {
if i.is_null() {
break;
}
if let Op::Store { dst, .. } = kd.kernel.ops[i].op {
if dst == *param {
stores_to_param.push(i);
}
}
i = kd.kernel.next_op(i);
}
debug_assert!(!stores_to_param.is_empty(), "store entry without store op in {depends_on:?}");
stores_to_param
})
.collect()
};
{
let kd = &self.kernels[depends_on];
for &tid in kd.outputs.iter().chain(kd.loads.iter()) {
if let TensorData::Eager { op_id, .. } | TensorData::Promoted { op_id, .. } = self.tensors[tid] {
if kd.kernel.ops.contains_id(op_id) {
keep_alive.push(op_id);
}
}
}
}
let mut loads = self.kernels[depends_on].loads.clone();
for ¶m in &dead_params {
while let Some((store_op, src)) = {
let kd = &self.kernels[depends_on];
let mut found = None;
let mut i = kd.kernel.head;
for _ in 0..100_000 {
if i.is_null() {
break;
}
if let Op::Store { dst, src, .. } = kd.kernel.ops[i].op {
if dst == param {
found = Some((i, src));
break;
}
}
i = kd.kernel.next_op(i);
}
found
} {
self.kernels[depends_on].kernel.remove_op(store_op);
loads = self.kernels[depends_on].kernel.remove_unused_chain(src, &keep_alive, &loads);
}
self.kernels[depends_on].kernel.remove_op(param);
}
let kd = &mut self.kernels[depends_on];
kd.stores.retain(|&t| t != x);
kd.loads = loads;
}
}
self.tensors.remove(x);
}
TensorData::Graph { graph_id, .. } => {
self.tensors.remove(x);
if !graph_id.is_null() {
self.graphs[graph_id].ref_count -= 1;
if self.graphs[graph_id].ref_count == 0 {
self.remove_dead_graph(graph_id);
}
}
}
TensorData::Eager { kernel_id, op_id, .. } => {
if !kernel_id.is_null() {
debug_assert!(!op_id.is_null());
debug_assert!(!self.kernels[kernel_id].stores.contains(&x));
self.kernels[kernel_id].outputs.remove(&x);
if self.kernels[kernel_id].outputs.is_empty() && self.kernels[kernel_id].stores.is_empty() {
for &tid in &self.kernels[kernel_id].loads {
if let TensorData::Eager { kernel_id: k, .. } | TensorData::Promoted { kernel_id: k, .. } =
&mut self.tensors[tid]
{
if *k == kernel_id {
*k = KernelId::NULL;
}
}
}
let mut loads = std::mem::take(&mut self.kernels[kernel_id].loads);
self.kernels.remove(kernel_id);
#[cfg(feature = "debug_tensor_op")]
eprintln!("KDROP {kernel_id:?} dying={x} all_loads={loads:?}");
loads.retain(|&id| id != x);
for tid in loads {
self.release(tid);
}
} else {
let (out_ops, loads) = {
let kd = &self.kernels[kernel_id];
let mut out_ops: Vec<OpId> = kd
.outputs
.iter()
.map(|&tid| match self.tensors[tid] {
TensorData::Eager { op_id, .. } | TensorData::Promoted { op_id, .. } => op_id,
ref t => panic!("kernel output tid {tid} has unexpected tensor data {t:?}"),
})
.collect();
for &tid in &kd.loads {
if let TensorData::Eager { op_id, .. } | TensorData::Promoted { op_id, .. } = self.tensors[tid] {
if kd.kernel.ops.contains_id(op_id) {
out_ops.push(op_id);
}
}
}
(out_ops, kd.loads.clone())
};
let new_loads = self.kernels[kernel_id].kernel.remove_unused_chain(op_id, &out_ops, &loads);
let pruned = loads_dropped_by_prune(&loads, &new_loads);
for load in pruned {
if load != x {
self.release(load);
}
}
self.kernels[kernel_id].loads = new_loads;
if self.kernels[kernel_id].outputs.is_empty() {
self.materialize_kernel(kernel_id).expect("materialization in tensor detach from kernel failed");
}
}
}
self.tensors.remove(x);
}
TensorData::Promoted { kernel_id, op_id, graph_id, .. } => {
if !kernel_id.is_null() {
debug_assert!(!op_id.is_null());
debug_assert!(!self.kernels[kernel_id].stores.contains(&x));
self.kernels[kernel_id].outputs.remove(&x);
if self.kernels[kernel_id].outputs.is_empty() && self.kernels[kernel_id].stores.is_empty() {
for &tid in &self.kernels[kernel_id].loads {
if let TensorData::Eager { kernel_id: k, .. } | TensorData::Promoted { kernel_id: k, .. } =
&mut self.tensors[tid]
{
if *k == kernel_id {
*k = KernelId::NULL;
}
}
}
let mut loads = std::mem::take(&mut self.kernels[kernel_id].loads);
self.kernels.remove(kernel_id);
#[cfg(feature = "debug_tensor_op")]
eprintln!("KDROP {kernel_id:?} dying={x} all_loads={loads:?}");
loads.retain(|&id| id != x);
for tid in loads {
self.release(tid);
}
} else {
let (out_ops, loads) = {
let kd = &self.kernels[kernel_id];
let mut out_ops: Vec<OpId> = kd
.outputs
.iter()
.map(|&tid| match self.tensors[tid] {
TensorData::Eager { op_id, .. } | TensorData::Promoted { op_id, .. } => op_id,
ref t => panic!("kernel output tid {tid} has unexpected tensor data {t:?}"),
})
.collect();
for &tid in &kd.loads {
if let TensorData::Eager { op_id, .. } | TensorData::Promoted { op_id, .. } = self.tensors[tid] {
if kd.kernel.ops.contains_id(op_id) {
out_ops.push(op_id);
}
}
}
(out_ops, kd.loads.clone())
};
let new_loads = self.kernels[kernel_id].kernel.remove_unused_chain(op_id, &out_ops, &loads);
let pruned = loads_dropped_by_prune(&loads, &new_loads);
for load in pruned {
if load != x {
self.release(load);
}
}
self.kernels[kernel_id].loads = new_loads;
if self.kernels[kernel_id].outputs.is_empty() {
self.materialize_kernel(kernel_id).expect("materialization in tensor detach from kernel failed");
}
}
}
self.tensors.remove(x);
if !graph_id.is_null() {
self.graphs[graph_id].ref_count -= 1;
if self.graphs[graph_id].ref_count == 0 {
self.remove_dead_graph(graph_id);
}
}
}
}
}
pub(crate) fn remove_dead_graph(&mut self, graph_id: GraphId) {
self.graphs.remove(graph_id);
}
pub(crate) fn verify_tensor_invariants(&self) {
if !cfg!(debug_assertions) {
return;
}
for (tid, td) in self.tensors.iter() {
match td {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => {
let (kernel_id, op_id) = (*kernel_id, *op_id);
assert!(!kernel_id.is_null(), "verify: kernel-backed tensor {tid} has NULL kernel_id");
assert!(self.kernels.contains_id(kernel_id), "verify: tensor {tid} points at deleted kernel {kernel_id:?}");
let kd = &self.kernels[kernel_id];
assert!(
kd.outputs.contains(&tid) || kd.loads.contains(&tid),
"verify: tensor {tid} has kernel_id {kernel_id:?} but is neither in its outputs nor loads"
);
assert!(!op_id.is_null(), "verify: tensor {tid} has NULL op_id with live kernel {kernel_id:?}");
assert!(
kd.kernel.ops.contains_id(op_id),
"verify: tensor {tid} op {op_id:?} is not in kernel {kernel_id:?}'s op slab"
);
let mut reachable = false;
let mut i = kd.kernel.head;
for _ in 0..100_000 {
if i.is_null() {
break;
}
if i == op_id {
reachable = true;
break;
}
i = kd.kernel.next_op(i);
}
assert!(reachable, "verify: tensor {tid} op {op_id:?} is not reachable from kernel {kernel_id:?}'s op list");
}
TensorData::Leaf { .. } => {}
TensorData::PendingLeaf { depends_on, .. } => {
let depends_on = *depends_on;
if !depends_on.is_null() {
assert!(
self.kernels.contains_id(depends_on) && self.kernels[depends_on].stores.contains(&tid),
"verify: pending Leaf {tid} points at depends_on {depends_on:?} which does not store it"
);
}
}
TensorData::GraphLeaf { .. } | TensorData::Graph { .. } | TensorData::Symbolic { .. } => {}
}
}
for (kid, kd) in self.kernels.iter() {
for &tid in &kd.outputs {
assert!(!kd.loads.contains(&tid), "verify: kernel {kid:?} both outputs and loads tid {tid}");
assert!(!kd.stores.contains(&tid), "verify: kernel {kid:?} both outputs and stores tid {tid}");
}
for &tid in &kd.loads {
assert!(!kd.stores.contains(&tid), "verify: kernel {kid:?} both loads and stores tid {tid}");
}
if !kd.outputs.is_empty() && matches!(&kd.kernel.ops[kd.kernel.tail].op, Op::Param { .. }) {
panic!("verify: kernel {kid:?} with outputs ends in a bare Param op (pure-load kernel — use a Leaf)");
}
}
}
pub fn assert_graph_inventory(&self, graph_id: GraphId) {
let live = self
.tensors
.iter()
.filter(|(_, td)| match td {
TensorData::Graph { graph_id: g, rc, .. } | TensorData::Promoted { graph_id: g, rc, .. } if *g == graph_id => {
*rc > 0
}
TensorData::Graph { .. } | TensorData::Promoted { .. } => false,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Symbolic { .. } => false,
})
.count();
assert_eq!(
live as u64, self.graphs[graph_id].ref_count,
"graph {graph_id:?} affiliation desync: {live} live affiliated tensors but ref_count = {}",
self.graphs[graph_id].ref_count
);
}
pub fn new_eager_tensor(&mut self, shape: TensorId, dtype: DType, buffer_id: Buffer) -> TensorId {
let shape_id = if shape == TensorId::NULL {
ExprId::NULL
} else {
match self.tensors[shape] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("new_eager_tensor: shape tid {shape} is not symbolic: {t:?}"),
}
};
let tid = self.tensors.push(TensorData::Leaf { shape_id, dtype, buffer: buffer_id, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!("rc::new_eager_tensor -> tid={tid} Leaf shape_id={shape_id} rc=1 (handle only)");
tid
}
pub(crate) fn new_kernel_from_leaf(&mut self, x: TensorId) -> (KernelId, OpId) {
let (shape_id, dtype) = match self.tensors[x] {
TensorData::Leaf { shape_id, dtype, .. }
| TensorData::GraphLeaf { shape_id, dtype, .. }
| TensorData::PendingLeaf { shape_id, dtype, .. } => (shape_id, dtype),
TensorData::Eager { .. } | TensorData::Graph { .. } | TensorData::Promoted { .. } | TensorData::Symbolic { .. } => {
unreachable!("new_kernel_from_leaf: {:?}", self.tensors[x])
}
};
let kernel_id = self.kernels.push(KernelData {
outputs: Set::default(),
loads: Vec::new(),
stores: Vec::new(),
kernel: Kernel::from_device_id(Dev::Auto, None),
});
let shape = self.replay_expr(kernel_id, shape_id);
let op_id = self.kernels[kernel_id].kernel.push_back(Op::Param { dtype, kind: ParamKind::Global, shape });
self.kernels[kernel_id].loads.push(x);
self.retain(x);
(kernel_id, op_id)
}
pub fn new_constant_tensor(&mut self, value: Constant) -> TensorId {
let expr = self.intern(Expr::Constant { value });
self.tensors.push(TensorData::Symbolic { expr, rc: 1 })
}
pub fn new_full(&mut self, shape: TensorId, value: Constant) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::new_full(shape={shape:?}, value={value:?})");
let x = self.new_constant_tensor(value);
if shape.is_null() {
return x;
}
let expanded = self.expand(x, shape).unwrap();
self.release(x);
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={expanded}, {:?}", self.tensors[expanded]);
expanded
}
pub fn new_variable_tensor<T: Scalar>(&mut self, x: T) -> TensorId {
let value = Constant::new(x);
let expr = self.intern(Expr::Variable { value });
self.tensors.push(TensorData::Symbolic { expr, rc: 1 })
}
pub fn new_host_tensor<T: Scalar>(&mut self, shape: TensorId, data: Box<[T]>) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::new_host_tensor(shape={shape:?})");
if data.len() == 1 && shape.is_null() {
let tid = self.new_constant_tensor(Constant::new(data[0]));
return Ok(tid);
}
let dtype = T::dtype();
let bytes = (data.len() * dtype.bit_size() as usize).div_ceil(8);
debug_assert_eq!(data.len() * std::mem::size_of::<T>(), bytes);
let alloc_bytes = bytes + dtype.bit_size() as usize / 8;
let free_bytes = Pool::Host.free_bytes();
if alloc_bytes as Dim > free_bytes {
return Err(ZyxError::AllocationError(
format!("Attempted to allocate {alloc_bytes} B on host, but it only has {free_bytes} B free").into(),
));
}
let buf_id = Pool::Host.allocate(alloc_bytes as Dim)?;
let buffer_id = Buffer { pool: Pool::Host, buffer_id: buf_id };
{
let dst = Pool::Host.buffer_ptr_mut(buf_id);
unsafe {
std::ptr::copy_nonoverlapping(data.as_ptr().cast::<u8>(), dst, bytes);
std::ptr::write_bytes(dst.add(bytes), 0, alloc_bytes - bytes);
}
}
let shape_id = match self.tensors[shape] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("new_host_tensor: shape tid {shape} is not symbolic: {t:?}"),
};
let tid = self.tensors.push(TensorData::Leaf { shape_id, dtype, buffer: buffer_id, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={tid}, shape={:?} dtype={}", self.shape(tid), self.dtype(tid));
Ok(tid)
}
pub fn new_disk_tensor(
&mut self,
shape: TensorId,
dtype: DType,
path: &Path,
offset_bytes: u64,
) -> Result<TensorId, ZyxError> {
let shape_id = match self.tensors[shape] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("new_disk_tensor: shape tid {shape} is not symbolic: {t:?}"),
};
let resolved = self.resolve_symbolic_dims(shape_id);
let bytes: Dim = ((resolved.iter().product::<Dim>() * dtype.bit_size() as Dim) + 7) / 8;
let buffer_id = Buffer { pool: Pool::Disk, buffer_id: Pool::Disk.disk_buffer_from_path(bytes, path, offset_bytes) };
let tid = self.tensors.push(TensorData::Leaf { shape_id, dtype, buffer: buffer_id, rc: 1 });
Ok(tid)
}
pub fn cast(&mut self, x: TensorId, dtype: DType) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::cast(x={x}, dtype={dtype:?})");
match self.tensors[x] {
TensorData::Symbolic { expr, .. } => match self.exprs[expr].clone() {
Expr::Constant { value } => self.new_constant_tensor(value.cast(dtype)),
Expr::Variable { .. }
| Expr::Cast { .. }
| Expr::Unary { .. }
| Expr::Binary { .. }
| Expr::Stack { .. }
| Expr::Stack2 { .. }
| Expr::Stack3 { .. }
| Expr::Stack4 { .. }
| Expr::Stack5 { .. } => {
let nested = self.intern(Expr::Cast { x: expr, dtype });
let tid = self.tensors.push(TensorData::Symbolic { expr: nested, rc: 1 });
self.retain(x);
tid
}
},
TensorData::Eager { kernel_id, op_id, shape_id, .. } => {
let op_id = self.kernels[kernel_id].kernel.cast(op_id, dtype);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::PendingLeaf { shape_id, .. } | TensorData::Leaf { shape_id, .. } => {
let (kernel_id, op_id) = self.new_kernel_from_leaf(x);
let op_id = self.kernels[kernel_id].kernel.cast(op_id, dtype);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::GraphLeaf { class_id, graph_id, shape_id, .. }
| TensorData::Graph { class_id, graph_id, shape_id, .. }
| TensorData::Promoted { class_id, graph_id, shape_id, .. } => {
self.assert_graph_alive(graph_id);
let (_, class_id) = self.push_node(graph_id, Node::Cast { x: class_id, dtype });
self.graphs[graph_id].ref_count += 1;
debug_assert!(!shape_id.is_null(), "cast: input graph tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
tid
}
}
}
pub fn bitcast(&mut self, x: TensorId, dtype: DType) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::bitcast(x={x}, dtype={dtype:?})");
debug_assert_eq!(self.dtype(x).bit_size(), dtype.bit_size(), "bitcast requires equal bit widths");
match self.tensors[x] {
TensorData::Symbolic { .. } => {
todo!("bitcast of pure-symbolic tensors")
}
TensorData::Eager { kernel_id, op_id, shape_id, .. } => {
let op_id = self.kernels[kernel_id].kernel.bitcast(op_id, dtype);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::PendingLeaf { shape_id, .. } | TensorData::Leaf { shape_id, .. } => {
let (kernel_id, op_id) = self.new_kernel_from_leaf(x);
let op_id = self.kernels[kernel_id].kernel.bitcast(op_id, dtype);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::GraphLeaf { class_id, graph_id, shape_id, .. }
| TensorData::Graph { class_id, graph_id, shape_id, .. }
| TensorData::Promoted { class_id, graph_id, shape_id, .. } => {
self.assert_graph_alive(graph_id);
let (_, class_id) = self.push_node(graph_id, Node::Bitcast { x: class_id, dtype });
self.graphs[graph_id].ref_count += 1;
debug_assert!(!shape_id.is_null(), "bitcast: input graph tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
tid
}
}
}
pub fn unary(&mut self, x: TensorId, uop: UOp) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::unary(x={x}, uop={uop:?})");
self.verify_tensor_invariants();
match self.tensors[x] {
TensorData::Symbolic { expr, .. } => match self.exprs[expr].clone() {
Expr::Constant { value } => self.new_constant_tensor(value.unary(uop)),
Expr::Variable { .. }
| Expr::Cast { .. }
| Expr::Unary { .. }
| Expr::Binary { .. }
| Expr::Stack { .. }
| Expr::Stack2 { .. }
| Expr::Stack3 { .. }
| Expr::Stack4 { .. }
| Expr::Stack5 { .. } => {
let root = self.intern(Expr::Unary { x: expr, uop });
self.tensors.push(TensorData::Symbolic { expr: root, rc: 1 })
}
},
TensorData::Eager { kernel_id, op_id, shape_id, dtype, .. } => {
let op_id = self.kernels[kernel_id].kernel.unary(op_id, uop);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::PendingLeaf { shape_id, dtype, .. } | TensorData::Leaf { shape_id, dtype, .. } => {
let (kernel_id, op_id) = self.new_kernel_from_leaf(x);
let op_id = self.kernels[kernel_id].kernel.unary(op_id, uop);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::GraphLeaf { class_id, graph_id, shape_id, dtype, .. }
| TensorData::Graph { class_id, graph_id, shape_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, shape_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let (_node_id, class_id) = self.push_node(graph_id, Node::Unary { x: class_id, uop });
self.graphs[graph_id].ref_count += 1;
debug_assert!(!shape_id.is_null(), "unary: input graph tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, nid={_node_id:?}, cid={class_id:?}");
tid
}
}
}
pub fn binary(&mut self, x: TensorId, y: TensorId, bop: BOp) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::binary(x={x}, y={y}, bop={bop:?})");
self.verify_tensor_invariants();
for tid in [x, y] {
if let TensorData::PendingLeaf { depends_on, .. } = self.tensors[tid] {
debug_assert!(!depends_on.is_null(), "binary: PendingLeaf {tid} with null depends_on");
let seen: Set<TensorId> = self.kernels[depends_on].outputs.iter().copied().collect();
for out in seen {
self.add_store(out)?;
}
}
}
let rx = self.resolve_shape(x).len();
let ry = self.resolve_shape(y).len();
if !(rx == 0 || ry == 0) {
debug_assert_eq!(
self.resolve_shape(x),
self.resolve_shape(y),
"binary operands must be broadcast to equal shapes before runtime.binary (broadcasting is performed upstream by Tensor::broadcast)"
);
}
let x_sym = matches!(self.tensors[x], TensorData::Symbolic { .. });
let y_sym = matches!(self.tensors[y], TensorData::Symbolic { .. });
if x_sym && y_sym {
let ex = match self.tensors[x] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("binary: symbolic operand tid {x} is not Symbolic: {t:?}"),
};
let ey = match self.tensors[y] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("binary: symbolic operand tid {y} is not Symbolic: {t:?}"),
};
let root = self.intern(Expr::Binary { x: ex, y: ey, bop });
let tid = self.tensors.push(TensorData::Symbolic { expr: root, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> symbolic: tid={tid}");
return Ok(tid);
}
fn result_shape(rt: &Runtime, a: TensorId, b: TensorId) -> ExprId {
let sa = match rt.tensors[a] {
TensorData::Eager { shape_id, .. }
| TensorData::Leaf { shape_id, .. }
| TensorData::PendingLeaf { shape_id, .. }
| TensorData::GraphLeaf { shape_id, .. }
| TensorData::Graph { shape_id, .. }
| TensorData::Promoted { shape_id, .. } => shape_id,
TensorData::Symbolic { .. } => ExprId::NULL,
};
let sb = match rt.tensors[b] {
TensorData::Eager { shape_id, .. }
| TensorData::Leaf { shape_id, .. }
| TensorData::PendingLeaf { shape_id, .. }
| TensorData::GraphLeaf { shape_id, .. }
| TensorData::Graph { shape_id, .. }
| TensorData::Promoted { shape_id, .. } => shape_id,
TensorData::Symbolic { .. } => ExprId::NULL,
};
if sa.is_null() { sb } else { sa }
}
let x_is_graph = self.is_graph(x);
let y_is_graph = self.is_graph(y);
if x_is_graph || y_is_graph {
let graph_id = if x_is_graph {
match self.tensors[x] {
TensorData::Graph { graph_id, .. }
| TensorData::Promoted { graph_id, .. }
| TensorData::GraphLeaf { graph_id, .. } => graph_id,
ref t => unreachable!("{t:?}"),
}
} else {
match self.tensors[y] {
TensorData::Graph { graph_id, .. }
| TensorData::Promoted { graph_id, .. }
| TensorData::GraphLeaf { graph_id, .. } => graph_id,
ref t => unreachable!("{t:?}"),
}
};
self.assert_graph_alive(graph_id);
if !x_is_graph && !x_sym {
self.promote_to_graph(x, graph_id)?;
}
if !y_is_graph && !y_sym {
self.promote_to_graph(y, graph_id)?;
}
let cx = match self.tensors[x] {
TensorData::Graph { class_id, .. }
| TensorData::Promoted { class_id, .. }
| TensorData::GraphLeaf { class_id, .. } => class_id,
TensorData::Symbolic { expr, .. } => match self.exprs[expr].clone() {
Expr::Constant { value } => self.push_const(graph_id, value),
ref e => todo!("promote symbolic scalar tid {x} ({e:?}) into a graph"),
},
ref t => unreachable!("unreachable after promote: {t:?}"),
};
let cy = match self.tensors[y] {
TensorData::Graph { class_id, .. }
| TensorData::Promoted { class_id, .. }
| TensorData::GraphLeaf { class_id, .. } => class_id,
TensorData::Symbolic { expr, .. } => match self.exprs[expr].clone() {
Expr::Constant { value } => self.push_const(graph_id, value),
ref e => todo!("promote symbolic scalar tid {y} ({e:?}) into a graph"),
},
ref t => unreachable!("unreachable after promote: {t:?}"),
};
let class_id = self.push_binary_node(graph_id, cx, cy, bop);
{
let shape_id = result_shape(self, x, y);
debug_assert!(!shape_id.is_null(), "binary: non-scalar graph operands {x}/{y} have no shape expression");
self.graphs[graph_id].ref_count += 1;
let dtype = if bop.returns_bool() { DType::Bool } else { self.dtype(x) };
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
Ok(tid)
}
} else if x_sym || y_sym {
debug_assert!(!x_sym || !y_sym);
let sym = if x_sym { x } else { y };
let data = if x_sym { y } else { x };
let shape_id = result_shape(self, x, y);
let (kid, data_op) = match self.tensors[data] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
ref t => panic!("binary: non-slab operand tid {data} is not an eager tensor: {t:?}"),
};
let sym_op = self.replay_symbolic_into_kernel(kid, sym);
let op_id = if x_sym {
self.kernels[kid].kernel.binary(sym_op, data_op, bop)
} else {
self.kernels[kid].kernel.binary(data_op, sym_op, bop)
};
let dtype = if bop.returns_bool() { DType::Bool } else { self.dtype(data) };
let tid = self.tensors.push(TensorData::Eager { kernel_id: kid, op_id, shape_id, dtype, rc: 1 });
self.kernels[kid].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kid:?}, op_id={op_id:?}");
Ok(tid)
} else {
let sx = self.resolve_shape_without_variables(x);
let sy = self.resolve_shape_without_variables(y);
if !sx.is_empty() && !sy.is_empty() && sx != sy {
return Err(ZyxError::shape_error(
format!(
"binary: cannot prove operand shapes are equal: {sx:?} vs {sy:?} — a symbolic dim must be the same dim tensor in both operands, or concrete in both"
)
.into(),
));
}
let shape_id = result_shape(self, x, y);
let (mut kid_x, mut op_id_x) = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(x),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
panic!("binary: operand tid {x} is not an eager tensor: {:?}", self.tensors[x])
}
};
let (mut kid_y, mut op_id_y) = match self.tensors[y] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(y),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
panic!("binary: operand tid {y} is not an eager tensor: {:?}", self.tensors[y])
}
};
let (kernel_id, op_id) = if kid_x == kid_y {
let op_id = self.kernels[kid_x].kernel.binary(op_id_x, op_id_y, bop);
(kid_x, op_id)
} else {
let x_stores = !self.kernels[kid_x].stores.is_empty();
let y_stores = !self.kernels[kid_y].stores.is_empty();
match (x_stores, y_stores) {
(true, true) => {
self.add_store(x)?;
self.add_store(y)?;
}
(true, false) => self.add_store(x)?,
(false, true) => self.add_store(y)?,
(false, false) => {}
}
(kid_x, op_id_x) = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(x),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
unreachable!("add_store turned operand into unexpected data: {:?}", self.tensors[x])
}
};
(kid_y, op_id_y) = match self.tensors[y] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(y),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
unreachable!("add_store turned operand into unexpected data: {:?}", self.tensors[y])
}
};
let swap = self.kernels[kid_y].kernel.is_reduce() && !self.kernels[kid_x].kernel.is_reduce();
let (keep_kid, merge_kid, keep_op, merge_op) = if swap {
(kid_y, kid_x, op_id_y, op_id_x)
} else {
(kid_x, kid_y, op_id_x, op_id_y)
};
let op_map = self.merge_kernel(keep_kid, merge_kid)?;
let op_id = if swap {
self.kernels[keep_kid].kernel.binary(op_map[&merge_op], keep_op, bop)
} else {
self.kernels[keep_kid].kernel.binary(keep_op, op_map[&merge_op], bop)
};
(keep_kid, op_id)
};
let dtype = if bop.returns_bool() { DType::Bool } else { self.dtype(x) };
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
Ok(tid)
}
}
pub fn device(&self, x: TensorId) -> Dev {
if let Some(buf_id) = self.leaf_buffer(x) {
return match buf_id.pool {
Pool::Host => Dev::C,
Pool::Disk => Dev::Auto,
pool => Dev::all().into_iter().find(|d| d.pool() == pool).unwrap_or_else(|| {
panic!(
"device: tensor {x} lives in pool {pool:?}, which has no devices attached. The backend operating this pool was likely configured out or never initialized."
)
}),
};
}
match self.tensors[x] {
TensorData::Eager { kernel_id, .. } | TensorData::Promoted { kernel_id, .. } => self.kernels[kernel_id].kernel.dev,
TensorData::PendingLeaf { old_buffer: Some(buf_id), .. } => match buf_id.pool {
Pool::Host => Dev::C,
Pool::Disk => Dev::Auto,
pool => Dev::all().into_iter().find(|d| d.pool() == pool).unwrap_or_else(|| {
panic!(
"device: tensor {x} lives in pool {pool:?}, which has no devices attached. The backend operating this pool was likely configured out or never initialized."
)
}),
},
TensorData::PendingLeaf { depends_on, .. } if !depends_on.is_null() => self.kernels[depends_on].kernel.dev,
TensorData::PendingLeaf { .. } => Dev::Auto,
TensorData::Leaf { .. } | TensorData::GraphLeaf { .. } | TensorData::Graph { .. } | TensorData::Symbolic { .. } => Dev::Auto,
}
}
#[allow(clippy::wrong_self_convention)] pub fn to_device(&mut self, x: TensorId, device: Dev) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::to_device(x={x}, device={device:?})");
let dst_pool = device.pool();
if let Some(buf_id) = self.leaf_buffer(x) {
if buf_id.pool == dst_pool {
self.retain(x);
return Ok(x);
}
}
match self.tensors[x] {
TensorData::Leaf { shape_id, dtype, .. } => {
let buf_id = self.leaf_buffer(x).expect("to_device: realized Leaf has no buffer");
if buf_id.pool == dst_pool {
self.retain(x);
return Ok(x);
}
let shape = self.resolve_shape(x);
let bytes = ((shape.iter().product::<Dim>() * dtype.bit_size() as Dim) + 7) / 8;
let alloc_bytes = bytes + dtype.bit_size() as Dim / 8;
let dst_buf = dst_pool.allocate(alloc_bytes)?;
let dst_id = Buffer { pool: dst_pool, buffer_id: dst_buf };
dst_pool.pool_to_pool(buf_id.pool, buf_id.buffer_id, dst_id.buffer_id)?;
debug_assert!(!shape_id.is_null(), "to_device: eager tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Leaf { shape_id, dtype, buffer: dst_id, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={tid} (cross-pool copy {buf_id:?} -> {dst_id:?})");
Ok(tid)
}
TensorData::PendingLeaf { depends_on: kernel_id, shape_id, dtype, .. }
| TensorData::Eager { kernel_id, shape_id, dtype, .. } => {
if !kernel_id.is_null() {
let outputs: Vec<TensorId> = self.kernels[kernel_id].outputs.iter().copied().collect();
for out in outputs {
self.add_store(out)?;
}
}
let buf_id = self.leaf_buffer(x).expect("to_device: tensor {x} was not materialized");
if buf_id.pool == dst_pool {
self.retain(x);
return Ok(x);
}
let shape = self.resolve_shape(x);
let bytes = ((shape.iter().product::<Dim>() * dtype.bit_size() as Dim) + 7) / 8;
let alloc_bytes = bytes + dtype.bit_size() as Dim / 8;
let dst_buf = dst_pool.allocate(alloc_bytes)?;
let dst_id = Buffer { pool: dst_pool, buffer_id: dst_buf };
dst_pool.pool_to_pool(buf_id.pool, buf_id.buffer_id, dst_id.buffer_id)?;
debug_assert!(!shape_id.is_null(), "to_device: eager tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Leaf { shape_id, dtype, buffer: dst_id, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={tid} (cross-pool copy {buf_id:?} -> {dst_id:?})");
Ok(tid)
}
TensorData::GraphLeaf { class_id, graph_id, shape_id, .. }
| TensorData::Graph { class_id, graph_id, shape_id, .. }
| TensorData::Promoted { class_id, graph_id, shape_id, .. } => {
assert!(!self.graphs[graph_id].dead, "tape scope has ended (tensor belongs to a dead tape scope");
let (_node_id, cid) = self.push_node(graph_id, Node::ToDevice { x: class_id, device, time: 0 });
self.graphs[graph_id].ref_count += 1;
debug_assert!(!shape_id.is_null(), "to_device: input graph tensor {x} has no shape expression");
let dtype = self.dtype(x);
let tid = self.tensors.push(TensorData::Graph { class_id: cid, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={tid}, nid={_node_id:?}, cid={cid:?}");
Ok(tid)
}
TensorData::Symbolic { .. } => {
self.retain(x);
Ok(x)
}
}
}
pub fn contiguous(&mut self, x: TensorId) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::contiguous(x={x})");
self.verify_tensor_invariants();
match self.tensors[x] {
TensorData::Symbolic { .. }
| TensorData::Leaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::PendingLeaf { .. } => {
self.retain(x);
Ok(x)
}
TensorData::Graph { class_id, graph_id, shape_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, shape_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let (_node_id, cid) = self.push_node(graph_id, Node::Contiguous { x: class_id });
self.graphs[graph_id].ref_count += 1;
debug_assert!(!shape_id.is_null(), "contiguous: input graph tensor {x} has no shape expression");
let tid = self.tensors.push(TensorData::Graph { class_id: cid, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={tid}, nid={_node_id:?}, cid={cid:?}");
Ok(tid)
}
TensorData::Eager { .. } => {
let cast_tid = self.cast(x, self.dtype(x));
self.add_store(cast_tid)?;
#[cfg(feature = "debug_tensor_op")]
println!(" -> tid={cast_tid} (cast shim stored)");
Ok(cast_tid)
}
}
}
pub fn reduce(&mut self, x: TensorId, mut axes: Vec<UAxis>, rop: BOp) -> Result<TensorId, ZyxError> {
self.verify_tensor_invariants();
let rank = self.shape(x).len();
debug_assert!(!axes.is_empty(), "reduce must specify at least one axis");
debug_assert!(axes.iter().all(|&a| (a as usize) < rank), "reduce axis {axes:?} out of bounds for rank {rank}");
debug_assert!(
axes.len() == axes.iter().collect::<std::collections::BTreeSet<_>>().len(),
"reduce axes must be unique: {axes:?}"
);
axes.sort_unstable();
match self.tensors[x] {
TensorData::Graph { class_id, graph_id, dtype, .. } | TensorData::Promoted { class_id, graph_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let mut dims = self.shape(x).to_vec();
debug_assert!(!dims.is_empty(), "reduce: input graph tensor {x} has no shape expression");
for axis in axes.iter().rev() {
dims.remove(*axis as usize);
}
let shape_id = if dims.is_empty() {
let one_const = self.new_constant_tensor(Constant::idx(1i64));
let stacked = self.stack(&[one_const])?;
self.release(one_const);
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("reduce: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
expr
} else {
let stacked = self.stack(&dims)?;
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("reduce: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
expr
};
let (_node_id, class_id) =
self.push_node(graph_id, Node::Reduce { x: class_id, rop, axes: axes.into_boxed_slice() });
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
Ok(tid)
}
TensorData::Eager { dtype, .. } | TensorData::Leaf { dtype, .. } => {
let mut cur = x;
let mut owns_cur = false;
let n_axes = axes.len();
axes.sort_unstable_by(|a, b| b.cmp(a));
let mut dims = self.shape(x);
for axis in axes {
let rank = self.resolve_shape(cur).len();
let permute_axes: Vec<UAxis> = (0..rank as UAxis).filter(|&i| i != axis).chain([axis]).collect();
let prev = cur;
let prev_owned = owns_cur;
cur = self.permute(cur, permute_axes);
if prev_owned {
self.release(prev);
}
let (kid, op_id) = self.duplicate_or_store(cur, false)?;
let dims_ops = self.kernels[kid].kernel.shape_ids(op_id);
debug_assert!(!dims_ops.is_empty(), "reduce of scalar");
let reduce_axis = *dims_ops.last().unwrap();
let op_id = self.kernels[kid].kernel.push_back(Op::Reduce { x: op_id, rop, reduce_axis });
let mut kept_dims = dims.clone();
kept_dims.remove(axis);
let shape_id = if kept_dims.is_empty() {
ExprId::NULL
} else {
let stacked = self.stack(&kept_dims)?;
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("reduce: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
expr
};
let tid = self.tensors.push(TensorData::Eager { kernel_id: kid, op_id, shape_id, dtype, rc: 1 });
dims = kept_dims;
debug_assert_eq!(self.kernels[kid].outputs.len(), 0, "input into reduce must have empty outputs");
self.kernels[kid].outputs.insert(tid);
self.release(cur);
owns_cur = true;
cur = tid;
}
if rank == n_axes {
let (kid, op_id) = match self.tensors[cur] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
ref t => unreachable!("{t:?}"),
};
let one_const = self.new_constant_tensor(Constant::idx(1i64));
let stacked = self.stack(&[one_const])?;
self.release(one_const);
let shape_id = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("reduce: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
let one = self.kernels[kid].kernel.const_idx(1);
let op_id = self.kernels[kid].kernel.reshape(op_id, one);
match &mut self.tensors[cur] {
TensorData::Eager { op_id: slot, shape_id: slot_shape, .. } => {
*slot = op_id;
*slot_shape = shape_id;
}
ref t => unreachable!("{t:?}"),
}
}
#[cfg(feature = "debug_tensor_op")]
println!(
" -> eager: tid={cur}, op_id={:?}",
match self.tensors[cur] {
TensorData::Eager { op_id, .. } => op_id,
ref t => unreachable!("{t:?}"),
}
);
Ok(cur)
}
ref t => todo!("reduce of pure-slab tensor {t:?}"),
}
}
pub(super) fn stack(&mut self, tensors: &[TensorId]) -> Result<TensorId, ZyxError> {
debug_assert!(!tensors.is_empty(), "stack: empty");
#[cfg(feature = "debug_tensor_op")]
println!("runtime::stack(tensors={tensors:?})");
let dtype = self.dtype(tensors[0]);
if tensors.iter().all(|&t| matches!(self.tensors[t], TensorData::Symbolic { .. })) {
if tensors.len() == 1 {
self.retain(tensors[0]);
#[cfg(feature = "debug_tensor_op")]
println!(" -> symbolic: tid={} (1d shape, no stack node)", tensors[0]);
return Ok(tensors[0]);
}
let exprs: Vec<ExprId> = tensors
.iter()
.map(|&t| match self.tensors[t] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("stack: operand tid is not symbolic: {t:?}"),
})
.collect();
let expr = match exprs.len() {
2 => self.intern(Expr::Stack2 { exprs: [exprs[0], exprs[1]] }),
3 => self.intern(Expr::Stack3 { exprs: [exprs[0], exprs[1], exprs[2]] }),
4 => self.intern(Expr::Stack4 { exprs: [exprs[0], exprs[1], exprs[2], exprs[3]] }),
5 => self.intern(Expr::Stack5 { exprs: [exprs[0], exprs[1], exprs[2], exprs[3], exprs[4]] }),
_ => self.intern(Expr::Stack { exprs: exprs.into_boxed_slice() }),
};
let tid = self.tensors.push(TensorData::Symbolic { expr, rc: 1 });
for &t in tensors {
self.retain(t);
}
#[cfg(feature = "debug_tensor_op")]
println!(" -> symbolic: tid={tid}");
return Ok(tid);
}
if tensors.iter().any(|&t| self.is_graph(t)) {
let graph_id = tensors
.iter()
.find(|&&t| self.is_graph(t))
.map(|&t| match self.tensors[t] {
TensorData::Graph { graph_id, .. } | TensorData::Promoted { graph_id, .. } => graph_id,
ref t => unreachable!("{t:?}"),
})
.unwrap();
self.assert_graph_alive(graph_id);
for &t in tensors {
let is_pure_const = match self.tensors[t] {
TensorData::Symbolic { expr, .. } => matches!(self.exprs[expr], Expr::Constant { .. }),
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Graph { .. }
| TensorData::Promoted { .. } => false,
};
if !self.is_graph(t) && !is_pure_const {
self.promote_to_graph(t, graph_id)?;
}
}
let mut ops = Vec::with_capacity(tensors.len());
for &t in tensors {
ops.push(match self.tensors[t] {
TensorData::Graph { class_id, .. } | TensorData::Promoted { class_id, .. } => class_id,
TensorData::Symbolic { expr, .. } => match self.exprs[expr].clone() {
Expr::Constant { value } => self.push_const(graph_id, value),
ref e => todo!("stack: promote symbolic scalar tid {t} ({e:?}) into a graph"),
},
ref t => todo!("stack: promote symbolic scalar tid {t:?} into a graph"),
});
}
let (_, class_id) = self.push_node(graph_id, Node::Stack { ops: ops.into_boxed_slice() });
{
let len_const = self.new_constant_tensor(Constant::idx(tensors.len() as i64));
let mut shape_dims = Vec::with_capacity(tensors.len() + 1);
shape_dims.push(len_const);
shape_dims.extend(self.shape(tensors[0]));
let stacked = self.stack(&shape_dims)?;
self.release(len_const);
let shape_id = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("stack: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
Ok(tid)
}
} else {
let keep_kid = match self.tensors[tensors[0]] {
TensorData::Eager { kernel_id, .. } => kernel_id,
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(tensors[0]).0,
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
panic!("stack: operand tid {} is not an eager tensor: {:?}", tensors[0], self.tensors[tensors[0]])
}
};
let mut ops = Vec::with_capacity(tensors.len());
for &t in tensors {
let (mut kid, mut op) = match self.tensors[t] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(t),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
panic!("stack: operand is not an eager tensor: {:?}", self.tensors[t])
}
};
if kid != keep_kid {
if !self.kernels[kid].stores.is_empty() {
self.add_store(t)?;
(kid, op) = match self.tensors[t] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(t),
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => unreachable!("{:?}", self.tensors[t]),
};
}
if kid != keep_kid {
let op_map = self.merge_kernel(keep_kid, kid)?;
op = op_map[&op];
}
}
ops.push(op);
}
let op_id = self.kernels[keep_kid].kernel.stack(&ops);
let len_const = self.new_constant_tensor(Constant::idx(tensors.len() as i64));
let mut shape_dims = Vec::with_capacity(tensors.len() + 1);
shape_dims.push(len_const);
shape_dims.extend(self.shape(tensors[0]));
let stacked = self.stack(&shape_dims)?;
self.release(len_const);
let shape_id = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("stack: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
let tid = self.tensors.push(TensorData::Eager { kernel_id: keep_kid, op_id, shape_id, dtype, rc: 1 });
self.kernels[keep_kid].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={keep_kid:?}, op_id={op_id:?}");
Ok(tid)
}
}
pub(super) fn reshape(&mut self, x: TensorId, shape_id: TensorId) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::reshape(x={x}, shape={shape_id:?})");
let shape_expr = if shape_id == TensorId::NULL {
ExprId::NULL
} else {
match self.tensors[shape_id] {
TensorData::Symbolic { expr, .. } => expr,
TensorData::Eager { shape_id, .. }
| TensorData::Leaf { shape_id, .. }
| TensorData::PendingLeaf { shape_id, .. }
| TensorData::GraphLeaf { shape_id, .. }
| TensorData::Graph { shape_id, .. }
| TensorData::Promoted { shape_id, .. } => shape_id,
}
};
debug_assert_eq!(
self.resolve_shape(x).iter().product::<Dim>(),
self.resolve_symbolic_dims(shape_expr).iter().product::<Dim>(),
"reshape element count mismatch"
);
let dtype = self.dtype(x);
if self.is_graph(x) || self.is_graph(shape_id) {
let graph_id = if self.is_graph(x) {
match self.tensors[x] {
TensorData::Graph { graph_id, .. }
| TensorData::GraphLeaf { graph_id, .. }
| TensorData::Promoted { graph_id, .. } => graph_id,
ref t => unreachable!("{t:?}"),
}
} else {
match self.tensors[shape_id] {
TensorData::Graph { graph_id, .. }
| TensorData::GraphLeaf { graph_id, .. }
| TensorData::Promoted { graph_id, .. } => graph_id,
ref t => unreachable!("{t:?}"),
}
};
self.assert_graph_alive(graph_id);
if !self.is_graph(x) {
self.promote_to_graph(x, graph_id)?;
}
let x_class = match self.tensors[x] {
TensorData::Graph { class_id, .. }
| TensorData::GraphLeaf { class_id, .. }
| TensorData::Promoted { class_id, .. } => class_id,
ref t => unreachable!("{t:?}"),
};
let shape_class = match self.tensors[shape_id] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "reshape: shape belongs to a different tape scope");
class_id
}
TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, shape_id),
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {
panic!("reshape: shape operand {shape_id} is a data tensor, not a symbolic shape")
}
};
let (_, class_id) = self.push_node(graph_id, Node::Reshape { x: x_class, shape: shape_class });
{
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id: shape_expr, dtype, rc: 1 });
Ok(tid)
}
} else {
if let Some(buf_id) = self.leaf_buffer(x) {
if !shape_expr.is_null() {}
let dtype = self.dtype(x);
self.retain(x);
buf_id.pool.retain(buf_id.buffer_id);
let tid = self.tensors.push(TensorData::Leaf { shape_id: shape_expr, dtype, buffer: buf_id, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid} (Leaf, shares buffer with x={x})");
return Ok(tid);
}
let (kernel_id, op_id) = self.duplicate_or_store(x, false)?;
debug_assert_eq!(
self.kernels[kernel_id].outputs.len(),
0,
"input into reshape must have empty outputs before the shape kernel is merged"
);
let shape_op = self.replay_symbolic_into_kernel(kernel_id, shape_id);
let op_id = self.kernels[kernel_id].kernel.reshape(op_id, shape_op);
if !shape_expr.is_null() {}
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id: shape_expr, dtype, rc: 1 });
debug_assert_eq!(self.kernels[kernel_id].outputs.contains(&tid), false);
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
Ok(tid)
}
}
pub fn expand(&mut self, x: TensorId, shape_id: TensorId) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::expand(x={x}, shape={shape_id:?})");
let shape_expr = if shape_id == TensorId::NULL {
ExprId::NULL
} else {
match self.tensors[shape_id] {
TensorData::Symbolic { expr, .. } => expr,
TensorData::Eager { shape_id, .. }
| TensorData::Leaf { shape_id, .. }
| TensorData::PendingLeaf { shape_id, .. }
| TensorData::GraphLeaf { shape_id, .. }
| TensorData::Graph { shape_id, .. }
| TensorData::Promoted { shape_id, .. } => shape_id,
}
};
let dtype = self.dtype(x);
let sh = self.resolve_shape(x);
let target = self.resolve_symbolic_dims(shape_expr);
debug_assert!(
sh.len() <= target.len(),
"expand: input rank {} > target rank {}: {:?} -> {:?}",
sh.len(),
target.len(),
sh,
target
);
for (old, new) in sh.iter().copied().rev().zip(target.iter().copied().rev()) {
debug_assert!(old == new || old == 1, "expand: incompatible dims: {old} vs {new} in {:?} -> {:?}", sh, target);
}
match self.tensors[x] {
TensorData::Graph { class_id: x_class, graph_id, .. }
| TensorData::GraphLeaf { class_id: x_class, graph_id, .. }
| TensorData::Promoted { class_id: x_class, graph_id, .. } => {
self.assert_graph_alive(graph_id);
let shape_class = match self.tensors[shape_id] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "expand: shape belongs to a different tape scope");
class_id
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, shape_id),
};
let (_, class_id) = self.push_node(graph_id, Node::Expand { x: x_class, shape: shape_class });
{
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id: shape_expr, dtype, rc: 1 });
Ok(tid)
}
}
TensorData::Symbolic { .. } => {
let kid = self.kernels.push(KernelData {
outputs: Set::default(),
loads: Vec::new(),
stores: Vec::new(),
kernel: Kernel::from_device_id(Dev::Auto, None),
});
let val_op = self.replay_symbolic_into_kernel(kid, x);
let shape_op = self.replay_symbolic_into_kernel(kid, shape_id);
let op_id = self.kernels[kid].kernel.expand(val_op, shape_op);
if !shape_expr.is_null() {}
let tid = self.tensors.push(TensorData::Eager { kernel_id: kid, op_id, shape_id: shape_expr, dtype, rc: 1 });
self.kernels[kid].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!("runtime::expand(x={x}) -> eager from slab: tid={tid}, kid={kid:?}, op_id={op_id:?}");
Ok(tid)
}
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {
let force_store = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => self.kernels[kernel_id].kernel.is_preceded_by_compute(op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => false,
TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => {
panic!("expand: operand tid {x} is not an eager tensor: {:?}", self.tensors[x])
}
};
let (kernel_id, op_id) = self.duplicate_or_store(x, force_store)?;
debug_assert_eq!(
self.kernels[kernel_id].outputs.len(),
0,
"input into expand must have empty outputs before the shape kernel is merged"
);
let shape_op = self.replay_symbolic_into_kernel(kernel_id, shape_id);
let op_id = self.kernels[kernel_id].kernel.expand(op_id, shape_op);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id: shape_expr, dtype, rc: 1 });
debug_assert_eq!(self.kernels[kernel_id].outputs.contains(&tid), false);
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
Ok(tid)
}
}
}
pub fn permute(&mut self, x: TensorId, axes: Vec<UAxis>) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::permute(x={x}, axes={axes:?})");
self.verify_tensor_invariants();
let sh = self.resolve_shape(x).to_vec();
debug_assert_eq!(axes.len(), sh.len(), "permute: axes length {} != rank {}", axes.len(), sh.len());
{
let mut sorted = axes.clone();
sorted.sort();
debug_assert!(
sorted.iter().copied().eq(0..sh.len() as UAxis),
"permute: axes not a valid permutation: {axes:?} for rank {}",
sh.len()
);
}
if axes.iter().copied().eq(0..sh.len() as UAxis) {
self.retain(x);
return x;
}
let shape_id = {
let dims = self.shape(x);
let permuted = crate::shape::permute(&dims, &axes);
if permuted.is_empty() {
ExprId::NULL
} else {
let stacked = self.stack(&permuted).expect("permute: failed to build shape stack");
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. } => {
panic!("permute: shape tid {stacked} is not symbolic: {:?}", self.tensors[stacked])
}
};
self.release(stacked);
expr
}
};
match self.tensors[x] {
TensorData::Graph { class_id, graph_id, dtype, .. }
| TensorData::GraphLeaf { class_id, graph_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let (_, class_id) = self.push_node(graph_id, Node::Permute { x: class_id, axes: axes.into_boxed_slice() });
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
tid
}
TensorData::Eager { dtype, .. } | TensorData::Leaf { dtype, .. } | TensorData::PendingLeaf { dtype, .. } => {
let (kernel_id, op_id) = self.duplicate_or_store(x, false).unwrap();
let op_id = self.kernels[kernel_id]
.kernel
.push_back(Op::Move { x: op_id, mop: Box::new(MoveOp::Permute { axes: axes.into() }) });
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
debug_assert_eq!(self.kernels[kernel_id].outputs.len(), 0, "input into permute must have empty outputs");
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::Symbolic { .. } => todo!("permute of symbolic scalar"),
}
}
pub fn pad_zeros(&mut self, x: TensorId, axis: UAxis, lp: TensorId, len: TensorId) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::pad_zeros(x={x}, axis={axis}, lp={lp}, len={len})");
self.verify_tensor_invariants();
let rank = self.resolve_shape(x).len();
debug_assert!((axis as usize) < rank, "pad_zeros axis {axis} out of bounds for rank {rank}");
debug_assert!(
self.resolve_shape(lp).is_empty() || self.resolve_shape(lp) == [1],
"pad_zeros lp must be scalar, got {:?}",
self.resolve_shape(lp)
);
debug_assert!(
self.resolve_shape(len).is_empty() || self.resolve_shape(len) == [1],
"pad_zeros len must be scalar, got {:?}",
self.resolve_shape(len)
);
debug_assert!(
self.dtype(lp).is_int() && self.dtype(len).is_int(),
"pad_zeros bounds must be integer-typed, got lp={:?} len={:?}",
self.dtype(lp),
self.dtype(len)
);
let shape_id = {
let mut dims = self.shape(x);
dims[axis as usize] = len;
self.retain(len);
let stacked = self.stack(&dims).expect("pad_zeros: failed to build shape stack");
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Graph { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Promoted { .. } => {
panic!("pad_zeros: shape tid {stacked} is not symbolic: {:?}", self.tensors[stacked])
}
};
self.release(stacked);
expr
};
match self.tensors[x] {
TensorData::Graph { class_id, graph_id, dtype, .. }
| TensorData::GraphLeaf { class_id, graph_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let lp_class = match self.tensors[lp] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "pad_zeros: lp belongs to a different tape scope");
class_id
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, lp),
};
let len_class = match self.tensors[len] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "pad_zeros: len belongs to a different tape scope");
class_id
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, len),
};
let (_, class_id) = self.push_node(graph_id, Node::Pad { x: class_id, axis, lp: lp_class, len: len_class });
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
self.graphs[graph_id].ref_count += 1;
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
tid
}
TensorData::Eager { dtype, .. } | TensorData::Leaf { dtype, .. } | TensorData::PendingLeaf { dtype, .. } => {
let len_const = self
.resolve_symbolic(len)
.expect("pad_zeros: eager-arm len bound must be a resolvable scalar")
.as_dim()
.expect("pad_zeros: len bound does not evaluate to an integer");
let grows = len_const > self.resolve_shape(x)[axis as usize];
let force_store = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => {
grows && self.kernels[kernel_id].kernel.is_preceded_by_compute(op_id)
}
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } | TensorData::Promoted { .. } => false,
TensorData::Graph { .. } | TensorData::GraphLeaf { .. } | TensorData::Symbolic { .. } => {
unreachable!("{:?}", self.tensors[x])
}
};
let (kernel_id, op_id) = self.duplicate_or_store(x, force_store).unwrap();
debug_assert_eq!(
self.kernels[kernel_id].outputs.len(),
0,
"input into pad must have empty outputs before the bound kernels are merged"
);
let lp_op = self.replay_symbolic_into_kernel(kernel_id, lp);
let len_op = self.replay_symbolic_into_kernel(kernel_id, len);
let op_id = self.kernels[kernel_id]
.kernel
.push_back(Op::Move { x: op_id, mop: Box::new(MoveOp::Pad { axis, lp: lp_op, len: len_op }) });
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::Symbolic { .. } => todo!("pad_zeros of symbolic scalar"),
}
}
pub fn narrow(&mut self, x: TensorId, axis: UAxis, start: TensorId, len: TensorId) -> TensorId {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::narrow(x={x}, axis={axis}, start={start}, len={len})");
self.verify_tensor_invariants();
debug_assert!(
self.dtype(start).is_int() && self.dtype(len).is_int(),
"narrow bounds must be integer-typed, got start={:?} len={:?}",
self.dtype(start),
self.dtype(len)
);
debug_assert!(
self.resolve_shape(start).is_empty() || self.resolve_shape(start) == [1],
"narrow start must be scalar, got {:?}",
self.resolve_shape(start)
);
debug_assert!(
self.resolve_shape(len).is_empty() || self.resolve_shape(len) == [1],
"narrow len must be scalar, got {:?}",
self.resolve_shape(len)
);
let sh = self.resolve_shape(x).to_vec();
debug_assert!(axis < sh.len() as UAxis, "narrow: axis {axis} out of range for rank {}", sh.len());
let shape_id = {
let mut dims = self.shape(x);
dims[axis as usize] = len;
self.retain(len);
let stacked = self.stack(&dims).expect("narrow: failed to build shape stack");
let expr = match self.tensors[stacked] {
TensorData::Symbolic { expr, .. } => expr,
ref t => panic!("narrow: shape tid {stacked} is not symbolic: {t:?}"),
};
self.release(stacked);
expr
};
match self.tensors[x] {
TensorData::Graph { class_id, graph_id, dtype, .. }
| TensorData::GraphLeaf { class_id, graph_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let start_class = match self.tensors[start] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "narrow: start belongs to a different tape scope");
class_id
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, start),
};
let len_class = match self.tensors[len] {
TensorData::Graph { class_id, graph_id: g, .. }
| TensorData::GraphLeaf { class_id, graph_id: g, .. }
| TensorData::Promoted { class_id, graph_id: g, .. } => {
assert!(g == graph_id, "narrow: len belongs to a different tape scope");
class_id
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => self.replay_symbolic_into_graph(graph_id, len),
};
let (_, class_id) =
self.push_node(graph_id, Node::Narrow { x: class_id, axis, start: start_class, len: len_class });
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
self.graphs[graph_id].ref_count += 1;
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
tid
}
TensorData::Eager { dtype, .. } | TensorData::Leaf { dtype, .. } | TensorData::PendingLeaf { dtype, .. } => {
let (kernel_id, op_id) = self.duplicate_or_store(x, false).unwrap();
debug_assert_eq!(
self.kernels[kernel_id].outputs.len(),
0,
"input into narrow must have empty outputs before the bound kernels are merged"
);
let start_op = self.replay_symbolic_into_kernel(kernel_id, start);
let len_op = self.replay_symbolic_into_kernel(kernel_id, len);
let op_id = self.kernels[kernel_id]
.kernel
.push_back(Op::Move { x: op_id, mop: Box::new(MoveOp::Narrow { axis, start: start_op, len: len_op }) });
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
tid
}
TensorData::Symbolic { .. } => todo!("narrow of symbolic scalar"),
}
}
pub fn flip(&mut self, x: TensorId, mut axes: Vec<UAxis>) -> Result<TensorId, ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::flip(x={x}, axes={axes:?})");
self.verify_tensor_invariants();
let sh = self.resolve_shape(x).to_vec();
if axes.is_empty() {
return Err(ZyxError::shape_error(format!("flip: axes must not be empty for tensor of shape {sh:?}").into()));
}
for &axis in &axes {
if axis >= sh.len() {
return Err(ZyxError::shape_error(format!("Axis {axis} is out of range of rank {}", sh.len()).into()));
}
}
axes.sort_unstable();
axes.dedup();
let shape_id = match self.tensors[x] {
TensorData::Eager { shape_id, .. }
| TensorData::Leaf { shape_id, .. }
| TensorData::Graph { shape_id, .. }
| TensorData::GraphLeaf { shape_id, .. }
| TensorData::Promoted { shape_id, .. } => shape_id,
ref t => todo!("flip of pure-slab tensor {t:?}"),
};
if shape_id != ExprId::NULL {}
match self.tensors[x] {
TensorData::Graph { class_id, graph_id, dtype, .. }
| TensorData::GraphLeaf { class_id, graph_id, dtype, .. }
| TensorData::Promoted { class_id, graph_id, dtype, .. } => {
self.assert_graph_alive(graph_id);
let (_, class_id) = self.push_node(graph_id, Node::Flip { x: class_id, axes: axes.into_boxed_slice() });
self.graphs[graph_id].ref_count += 1;
let tid = self.tensors.push(TensorData::Graph { class_id, graph_id, shape_id, dtype, rc: 1 });
#[cfg(feature = "debug_tensor_op")]
println!(" -> graph: tid={tid}, graph_id={graph_id:?}, class_id={class_id:?}");
Ok(tid)
}
TensorData::Eager { dtype, .. } | TensorData::Leaf { dtype, .. } => {
let (kernel_id, op_id) = self.duplicate_or_store(x, false).unwrap();
let op_id = self.kernels[kernel_id].kernel.flip(op_id, &axes);
let tid = self.tensors.push(TensorData::Eager { kernel_id, op_id, shape_id, dtype, rc: 1 });
debug_assert_eq!(self.kernels[kernel_id].outputs.len(), 0, "input into flip must have empty outputs");
self.kernels[kernel_id].outputs.insert(tid);
#[cfg(feature = "debug_tensor_op")]
println!(" -> eager: tid={tid}, kid={kernel_id:?}, op_id={op_id:?}");
Ok(tid)
}
ref t => unreachable!("shape extraction already rejected non-slab shapes: {t:?}"),
}
}
pub fn load<T: Scalar>(&mut self, x: TensorId, data: &mut [T]) -> Result<(), ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::load(x={x})");
self.verify_tensor_invariants();
if self.leaf_buffer(x).is_none() {
if let Some(c) = self.resolve_symbolic(x) {
let v = match c.cast(T::dtype()) {
Constant::BF16(v) => T::from_bf16(bf16::from_le_bytes(v)),
Constant::F16(v) => T::from_f16(f16::from_le_bytes(v)),
Constant::F32(v) => T::from_f32(f32::from_le_bytes(v)),
Constant::F64(v) => T::from_f64(f64::from_le_bytes(v)),
Constant::F8E4M3(v) => T::from_f32(f8e4m3::from_bits(v).to_f32()),
Constant::F8E5M2(v) => T::from_f32(f8e5m2::from_bits(v).to_f32()),
Constant::U8(v) => T::from_u8(v),
Constant::U16(v) => T::from_u16(v),
Constant::U32(v) => T::from_u32(v),
Constant::U64(v) => T::from_u64(u64::from_le_bytes(v)),
Constant::I8(v) => T::from_i8(v),
Constant::I16(v) => T::from_i16(v),
Constant::I32(v) => T::from_i32(v),
Constant::I64(v) => T::from_i64(i64::from_le_bytes(v)),
Constant::Bool(v) => T::from_bool(v),
};
for d in data.iter_mut() {
*d = v;
}
return Ok(());
}
}
let dt = self.dtype(x);
if dt != T::dtype() {
return Err(ZyxError::DTypeError(format!("loading dtype {}, but the data has dtype {dt}", T::dtype()).into()));
}
let shape_numel: Dim = self.resolve_shape(x).iter().product();
if (data.len() as Dim) > shape_numel {
return Err(ZyxError::AllocationError(
format!("load buffer of {} elements is larger than tensor with {shape_numel} elements", data.len()).into(),
));
}
if let Some(value) = self.resolve_symbolic(x) {
let bytes = (data.len() * T::bit_size() as usize).div_ceil(8);
let byte_slice = unsafe { std::slice::from_raw_parts_mut(data.as_mut_ptr().cast(), bytes) };
let value_bytes = value.to_le_bytes();
byte_slice[..value_bytes.len()].copy_from_slice(&value_bytes);
return Ok(());
}
if let TensorData::PendingLeaf { depends_on, .. } = self.tensors[x] {
let depends_on = depends_on;
if !depends_on.is_null() && self.kernels.contains_id(depends_on) {
let outputs: Set<TensorId> = self.kernels[depends_on].outputs.iter().copied().collect();
for tid in outputs {
self.add_store(tid)?;
}
if self.kernels.contains_id(depends_on)
&& self.kernels[depends_on].outputs.is_empty()
&& !self.kernels[depends_on].stores.is_empty()
{
self.materialize_kernel(depends_on)?;
}
}
}
let Some(mut buffer_id) = self.leaf_buffer(x) else {
let this = &mut *self;
let pending = match this.tensors[x] {
TensorData::PendingLeaf { depends_on, .. } => depends_on,
TensorData::Eager { .. } => KernelId::NULL,
TensorData::Graph { .. } => return Err(ZyxError::graph_tensor_not_realized(x)),
TensorData::Leaf { .. } | TensorData::GraphLeaf { .. } => {
unreachable!("load: buffered tensor {x} has no buffer: {:?}", self.tensors[x])
}
TensorData::Promoted { .. } | TensorData::Symbolic { .. } => {
panic!("load: tensor {x} has no buffer and cannot be materialized: {:?}", self.tensors[x])
}
};
if !pending.is_null() {
let outputs: Set<TensorId> = this.kernels[pending].outputs.iter().copied().collect();
for tid in outputs {
this.add_store(tid)?;
}
}
if this.leaf_buffer(x).is_none() {
let kid = match this.tensors[x] {
TensorData::Eager { kernel_id, .. } => kernel_id,
TensorData::Graph { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {
return Err(ZyxError::graph_tensor_not_realized(x));
}
ref t => panic!("load: tensor {x} has no buffer and cannot be materialized: {t:?}"),
};
let seen: Set<TensorId> = this.kernels[kid].outputs.iter().copied().collect();
for tid in seen {
this.add_store(tid)?;
}
}
let bytes = (data.len() * T::bit_size() as usize).div_ceil(8);
let byte_slice = unsafe { std::slice::from_raw_parts_mut(data.as_mut_ptr().cast(), bytes) };
let buffer_id = this.leaf_buffer(x).expect("load: tensor has no buffer after materialization");
buffer_id.pool.pool_to_host(buffer_id.buffer_id, byte_slice)?;
#[cfg(feature = "debug_tensor_op")]
println!(" -> x={x}, {:?}", self.tensors[x]);
return Ok(());
};
if let TensorData::PendingLeaf { depends_on, .. } = self.tensors[x] {
debug_assert!(!depends_on.is_null(), "load: PendingLeaf {x} with null depends_on");
if !depends_on.is_null() {
let seen: Set<TensorId> = self.kernels[depends_on].outputs.iter().copied().collect();
for tid in seen {
self.add_store(tid)?;
}
buffer_id = self.leaf_buffer(x).ok_or_else(|| {
ZyxError::AllocationError(format!("load: tensor {x} lost its buffer during pending store").into())
})?;
}
}
let bytes = (data.len() * T::bit_size() as usize).div_ceil(8);
let byte_slice = unsafe { std::slice::from_raw_parts_mut(data.as_mut_ptr().cast(), bytes) };
buffer_id.pool.pool_to_host(buffer_id.buffer_id, byte_slice)?;
#[cfg(feature = "debug_tensor_op")]
println!(" -> x={x}, {:?}", self.tensors[x]);
Ok(())
}
pub fn assign(&mut self, dst: TensorId, src: TensorId) -> Result<(), ZyxError> {
#[cfg(feature = "debug_tensor_op")]
println!("runtime::assign(dst={dst}, src={src})");
if src == dst {
return Err(ZyxError::shape_error(format!("assign: src and dst are the same tensor {dst}").into()));
}
self.verify_tensor_invariants();
let dst_dtype = self.dtype(dst);
let src_dtype = self.dtype(src);
if dst_dtype != src_dtype {
return Err(ZyxError::DTypeError(format!("assign dtype mismatch: dst={dst_dtype}, src={src_dtype}").into()));
}
let dst_shape = self.resolve_shape(dst);
let src_shape = self.resolve_shape(src);
if dst_shape != src_shape {
return Err(ZyxError::shape_error(format!("assign shape mismatch: dst={dst_shape:?}, src={src_shape:?}").into()));
}
let dst_syms = self.resolve_shape_without_variables(dst);
let src_syms = self.resolve_shape_without_variables(src);
if !dst_syms.is_empty() && !src_syms.is_empty() && dst_syms != src_syms {
return Err(ZyxError::shape_error(
format!(
"assign: cannot prove dst and src shapes are equal: {dst_syms:?} vs {src_syms:?} — a symbolic dim must be the same dim tensor in both operands, or concrete in both"
)
.into(),
));
}
match self.tensors[dst] {
TensorData::Graph { class_id: dst_cid, graph_id, .. }
| TensorData::GraphLeaf { class_id: dst_cid, graph_id, .. }
| TensorData::Promoted { class_id: dst_cid, graph_id, .. } => {
if dst == src {
return Err(ZyxError::ShapeError("assign: dst equals src (self-assign)".into()));
}
self.assert_graph_alive(graph_id);
match self.tensors[src] {
TensorData::Graph { graph_id: g, .. }
| TensorData::GraphLeaf { graph_id: g, .. }
| TensorData::Promoted { graph_id: g, .. } => {
if g != graph_id {
panic!("tensor belongs to a different tape scope");
}
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => {
self.promote_to_graph(src, graph_id)?;
}
}
let mut dst_leaf_cid = dst_cid;
let graph = &self.graphs[graph_id];
loop {
match graph.nodes[dst_leaf_cid].node {
Node::Pad { x, .. }
| Node::Flip { x, .. }
| Node::Expand { x, .. }
| Node::Reshape { x, .. }
| Node::Narrow { x, .. }
| Node::Permute { x, .. } => dst_leaf_cid = x,
Node::After { .. } | Node::Leaf { .. } => break,
ref op => unreachable!("{op:?}"),
}
}
let mut leaf_cid = dst_leaf_cid;
while let Node::After { x, .. } = &graph.nodes[leaf_cid].node {
leaf_cid = *x;
}
let dst_leaf = graph.leaf_map[&leaf_cid];
let src_cid = match self.tensors[src] {
TensorData::Graph { class_id, .. }
| TensorData::GraphLeaf { class_id, .. }
| TensorData::Promoted { class_id, .. } => class_id,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => {
unreachable!("{:?}", self.tensors[src])
}
};
let (_node_id, assign_cid) = self.push_node(graph_id, Node::Assign { dst: dst_cid, src: src_cid });
let leaf_class = self.push_node(graph_id, Node::After { x: dst_leaf_cid, dep: assign_cid }).1;
let dst_class = self.push_node(graph_id, Node::After { x: dst_cid, dep: assign_cid }).1;
for (tid, class_id) in [(dst_leaf, leaf_class), (dst, dst_class)] {
match &mut self.tensors[tid] {
TensorData::Graph { class_id: c, .. }
| TensorData::GraphLeaf { class_id: c, .. }
| TensorData::Promoted { class_id: c, .. } => *c = class_id,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::Symbolic { .. } => {
panic!("assign: tensor {tid} has no graph class to re-point: {:?}", self.tensors[tid])
}
}
}
#[cfg(feature = "debug_tensor_op")]
println!(" -> assign_cid={assign_cid:?}");
return Ok(());
}
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } | TensorData::Symbolic { .. } => {
}
}
let (src_kid, src_op) = match self.tensors[src] {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => self.new_kernel_from_leaf(src),
TensorData::Graph { .. } | TensorData::GraphLeaf { .. } | TensorData::Symbolic { .. } => {
panic!("assign: src {src} is not an eager/promoted tensor: {:?}", self.tensors[src])
}
};
if let TensorData::PendingLeaf { depends_on, .. } = self.tensors[dst] {
debug_assert!(!depends_on.is_null(), "assign: PendingLeaf {dst} with null depends_on");
let seen: Set<TensorId> = self.kernels[depends_on].outputs.iter().copied().collect();
for tid in seen {
self.add_store(tid)?;
}
}
if let TensorData::Leaf { shape_id: dst_shape_id, buffer: dst_buf, rc: dst_rc, .. } = self.tensors[dst] {
let dtype = self.dtype(dst);
let (kernel_id, src_op) = match self.tensors[src] {
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {
let (kid, op) = self.new_kernel_from_leaf(src);
(kid, op)
}
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Graph { .. } | TensorData::GraphLeaf { .. } | TensorData::Symbolic { .. } => {
panic!("assign: src {src} is not an eager/promoted tensor: {:?}", self.tensors[src])
}
};
let dst_shape_op = self.replay_expr(kernel_id, dst_shape_id);
let mut_param =
self.kernels[kernel_id].kernel.push_back(Op::Param { dtype, kind: ParamKind::GlobalMut, shape: dst_shape_op });
self.kernels[kernel_id].kernel.store(mut_param, src_op, OpId::NULL);
self.kernels[kernel_id].stores.push(dst);
self.tensors[dst] = TensorData::PendingLeaf {
old_buffer: Some(dst_buf),
depends_on: kernel_id,
shape_id: dst_shape_id,
dtype,
rc: dst_rc,
};
return Ok(());
}
let (dst_kid, dst_op) = match self.tensors[dst] {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
ref t => panic!("assign: dst {dst} is not an eager/promoted tensor: {t:?}"),
};
if self.kernels[dst_kid].outputs.iter().any(|&e| e != dst) {
return Err(ZyxError::ShapeError(
format!("assign: dst kernel {dst_kid:?} has other outputs {:?}, only dst allowed", self.kernels[dst_kid].outputs)
.into(),
));
}
for op in self.kernels[dst_kid].kernel.ops.values() {
if !matches!(op.op, Op::Param { .. } | Op::Move { .. } | Op::Const(_) | Op::Stack { .. }) {
return Err(ZyxError::ShapeError(
format!("assign: dst kernel {dst_kid:?} has unsupported op {:?}, only movement ops allowed", op.op).into(),
));
}
}
if src_kid == dst_kid {
return Err(ZyxError::ShapeError(
format!("assign: src and dst share kernel {dst_kid:?}; dst must be a separate movement-only kernel").into(),
));
}
if !self.kernels[dst_kid].stores.is_empty() {
return Err(ZyxError::ShapeError(
format!("assign: dst kernel {dst_kid:?} has stores {}; expected none", self.kernels[dst_kid].stores.len()).into(),
));
}
if self.kernels[src_kid].loads.contains(&dst) {
return Err(ZyxError::ShapeError(
format!("assign: src kernel {dst_kid:?} loads dst tensor, not allowed to avoid data races").into(),
));
}
let dst_kernel_loads = self.kernels[dst_kid].loads.clone();
let dst_org = {
let mut buffer_loads =
dst_kernel_loads.iter().copied().filter(|&t| {
!matches!(self.tensors[t], TensorData::Symbolic { expr, .. } if matches!(self.exprs[expr], Expr::Variable { .. }))
});
match (buffer_loads.next(), buffer_loads.next()) {
(Some(t), None) => t,
(None, _) => {
return Err(ZyxError::ShapeError(
"assign: dst kernel has no backing buffer; its base was never materialized \
into pool storage — call `.contiguous()` on it before assign"
.into(),
));
}
(Some(_), Some(_)) => return Err(ZyxError::ShapeError(
"assign: dst kernel contains more than one buffer load; dst must be a movement-only view of exactly one base"
.into(),
)),
}
};
match self.tensors[dst_org] {
TensorData::PendingLeaf { depends_on, .. } if !depends_on.is_null() => {
assert!(
depends_on != src_kid,
"assign: dst base {dst_org} is pending on src's kernel {src_kid:?}; assign would interleave with its own store"
);
}
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {}
ref t => panic!("assign: dst base {dst_org} in unexpected state {t:?}"),
}
let kernel = self.kernels[dst_kid].kernel.clone();
let loads = self.kernels[dst_kid].loads.clone();
for t in &loads {
assert!(
*t == dst_org
|| matches!(self.tensors[*t], TensorData::Symbolic { expr, .. } if matches!(self.exprs[expr], Expr::Variable { .. })),
"assign: dst kernel load {t} is neither the buffer nor a known variable"
);
}
{
let mut n_params = 0usize;
let mut p = kernel.head;
while !p.is_null() {
if matches!(&kernel.ops[p].op, Op::Param { .. }) {
n_params += 1;
}
p = kernel.next_op(p);
}
assert_eq!(n_params, loads.len(), "assign: dst kernel param/loads count mismatch");
}
let mut dst_param = dst_op;
for _ in 0..100 {
match kernel.ops[dst_param].op {
Op::Move { x, .. } => {
dst_param = x;
}
Op::Param { .. } => {
break;
}
ref op => {
return Err(ZyxError::ShapeError(
format!(
"assign: dst movement chain contains a non-Move op {op:?}; the dst base must bottom out in a Param"
)
.into(),
));
}
}
}
if !matches!(kernel.ops[dst_param].op, Op::Param { .. }) {
return Err(ZyxError::ShapeError(
"assign: dst movement chain exceeds 100 Move ops; the dst base does not bottom out in a Param".into(),
));
}
let mut op_map = Map::default();
let mut new_def_loads: Vec<TensorId> = Vec::new();
let mut required: Set<OpId> = Set::default();
{
let mut stack: Vec<OpId> = Vec::new();
let mut oid = kernel.head;
while !oid.is_null() {
stack.push(oid);
oid = kernel.next_op(oid);
}
while let Some(id) = stack.pop() {
if !required.insert(id) {
continue;
}
stack.extend(kernel.ops[id].op.parameters());
}
debug_assert!(stack.is_empty(), "assign replay: dependency walk did not finish");
}
let mut def_i = 0usize;
let mut op_id = kernel.head;
while !op_id.is_null() {
if required.contains(&op_id) {
let mut op = kernel.ops[op_id].op.clone();
if let Op::Move { x, .. } = &mut op {
if op_map.get(x).is_none() {
*x = op_map[&dst_param];
}
}
for p in op.parameters_mut() {
*p =
op_map.get(p).copied().expect("assign replay: dependency was not copied before its user despite closure");
}
let mut new_def_load: Option<TensorId> = None;
if let Op::Param { kind, .. } = &mut op {
if op_id == dst_param {
*kind = ParamKind::GlobalMut;
}
assert!(
matches!(kind, ParamKind::GlobalMut | ParamKind::Variable),
"assign: unexpected param kind {kind:?} in dst movement kernel"
);
if *kind != ParamKind::GlobalMut {
new_def_load = Some(loads[def_i]);
}
def_i += 1;
}
let new_id = self.kernels[src_kid].kernel.push_back(op);
if let Some(load) = new_def_load {
new_def_loads.push(load);
}
op_map.insert(op_id, new_id);
}
op_id = kernel.next_op(op_id);
}
let dst_op = op_map.get(&dst_op).copied().unwrap_or(op_map[&dst_param]);
self.kernels[src_kid].kernel.store(dst_op, src_op, OpId::NULL);
debug_assert!(
matches!(self.tensors[dst_org], TensorData::Leaf { .. } | TensorData::PendingLeaf { .. }),
"assign: dst base {dst_org} is not a leaf/pending leaf"
);
self.kernels[src_kid].stores.push(dst_org);
for load in new_def_loads {
debug_assert!(
matches!(
self.tensors[load],
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } | TensorData::Symbolic { .. }
),
"assign: replayed load {load} is not a leaf/pending leaf/symbolic"
);
self.kernels[src_kid].loads.push(load);
self.retain(load);
}
#[cfg(debug_assertions)]
{
let kd = &self.kernels[src_kid];
assert!(!kd.loads.contains(&dst_org), "assign: GlobalMut store target {dst_org} leaked into loads");
let mut n_non_mut = 0usize;
let mut p = kd.kernel.head;
while !p.is_null() {
if let Op::Param { kind, .. } = &kd.kernel.ops[p].op {
if *kind != ParamKind::GlobalMut {
n_non_mut += 1;
}
}
p = kd.kernel.next_op(p);
}
assert_eq!(n_non_mut, kd.loads.len(), "assign: loads/defines alignment broken for kernel {:?}", src_kid);
}
match self.tensors[dst_org] {
TensorData::PendingLeaf { depends_on, .. } if !depends_on.is_null() => {
for out in self.kernels[depends_on].outputs.clone() {
self.add_store(out)?;
}
}
TensorData::Eager { .. } | TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {}
ref t => panic!("assign: dst base {dst_org} in unexpected state {t:?}"),
}
debug_assert!(self.leaf_buffer(dst_org).is_some(), "assign: dst base {dst_org} has no buffer after materialization");
let (dst_shape_id, dst_dtype, dst_rc, dst_buf) = match self.tensors[dst_org] {
TensorData::Leaf { shape_id, dtype, rc, buffer } => (shape_id, dtype, rc, buffer),
ref t => panic!("assign: dst base {dst_org} is not a realized leaf: {t:?}"),
};
self.tensors[dst_org] = TensorData::PendingLeaf {
old_buffer: Some(dst_buf),
depends_on: src_kid,
shape_id: dst_shape_id,
dtype: dst_dtype,
rc: dst_rc,
};
let outputs: Vec<TensorId> = self.kernels[src_kid].outputs.iter().copied().collect();
if outputs.is_empty() {
self.materialize_kernel(src_kid)?;
} else {
for tid in outputs {
self.add_store(tid)?;
}
}
Ok(())
}
#[allow(unused)]
pub fn deinitialize(&mut self) {
#[cfg(feature = "time")]
{
let lock = crate::ET.lock();
let mut timings: Vec<_> = lock.iter().map(|(name, &(total_us, count))| (name.clone(), total_us, count)).collect();
timings.sort_by_key(|a| std::cmp::Reverse(a.1));
println!("\n=== Timing Info (sorted by total time, descending) ===");
for (name, total_us, count) in timings {
let per_call = total_us.checked_div(count).unwrap_or(0);
println!("{name}: {total_us}us total, {per_call}us/call ({count} calls)");
}
}
self.tensors = Slab::new();
self.kernels = Slab::new();
}
pub const fn manual_seed(&mut self, seed: u64) {
self.rng = Rng::seed_from_u64(seed);
}
pub fn free_memory(&mut self) -> Dim {
Dev::all().iter().map(|d| d.pool().free_bytes()).max().unwrap_or(0)
}
}
impl Runtime {
fn duplicate_or_store(&mut self, x: TensorId, force_store: bool) -> Result<(KernelId, OpId), ZyxError> {
let (mut kid, mut op_id) = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } => {
return Ok(self.new_kernel_from_leaf(x));
}
TensorData::Graph { .. } | TensorData::GraphLeaf { .. } | TensorData::Symbolic { .. } => {
panic!("duplicate_or_store: tensor {x} is not an eager/promoted tensor: {:?}", self.tensors[x])
}
};
let contains_stores = self.kernels[kid].kernel.contains_stores();
let preceded_by_reduce = self.kernels[kid].kernel.is_preceded_by_reduce(op_id);
if force_store || contains_stores || preceded_by_reduce {
self.add_store(x)?;
(kid, op_id) = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. } | TensorData::PendingLeaf { .. } | TensorData::GraphLeaf { .. } => {
return Ok(self.new_kernel_from_leaf(x));
}
TensorData::Graph { .. } | TensorData::Symbolic { .. } => {
panic!("duplicate_or_store: tensor {x} is not an eager/promoted tensor: {:?}", self.tensors[x])
}
};
}
debug_assert!(self.kernels[kid].stores.is_empty(), "duplicated kernel must not have stores");
let old_loads = self.kernels[kid].loads.clone();
let (kernel, op_id, new_loads) = self.kernels[kid].kernel.duplicate_subkernel(op_id, &old_loads);
for &tid in &new_loads {
self.retain(tid);
}
kid = self.kernels.push(KernelData { outputs: Set::default(), loads: new_loads, stores: Vec::new(), kernel });
Ok((kid, op_id))
}
fn merge_kernel(&mut self, keep_kid: KernelId, merge_kid: KernelId) -> Result<Map<OpId, OpId>, ZyxError> {
debug_assert_ne!(keep_kid, merge_kid, "merge_kernel: cannot merge a kernel into itself");
debug_assert!(
self.kernels[merge_kid].stores.is_empty(),
"merge_kernel: merge kernel {merge_kid:?} has stores; add_store them before merging"
);
let KernelData { outputs: merge_outputs, loads: merge_loads, stores: merge_stores, kernel } =
unsafe { self.kernels.remove_and_return(merge_kid) };
let Kernel { ops: merge_ops, head: merge_head, .. } = kernel;
let mut const_map: Map<Constant, OpId> = Map::with_hasher(BuildHasherDefault::new());
{
let keep = &self.kernels[keep_kid].kernel;
let mut i = keep.head;
while !i.is_null() {
if let Op::Const(c) = keep.ops[i].op {
const_map.insert(c, i);
}
i = keep.ops[i].next;
}
}
let mut op_map: Map<OpId, OpId> = Map::with_hasher(BuildHasherDefault::new());
let mut i = merge_head;
while !i.is_null() {
let mut op = merge_ops[i].op.clone();
for param in op.parameters_mut() {
if let Some(&new_param) = op_map.get(param) {
*param = new_param;
}
}
if let Op::Const(c) = op {
if let Some(&survivor) = const_map.get(&c) {
op_map.insert(i, survivor);
i = merge_ops[i].next;
continue;
}
}
let new_op_id = self.kernels[keep_kid].kernel.push_back(op);
if let Op::Const(c) = merge_ops[i].op {
const_map.insert(c, new_op_id);
}
op_map.insert(i, new_op_id);
i = merge_ops[i].next;
}
for (_tid, t_data) in self.tensors.iter_mut() {
let (kernel_id, op_id) = match t_data {
TensorData::Eager { kernel_id, op_id, .. } | TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id),
TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Graph { .. }
| TensorData::Symbolic { .. } => continue,
};
if *kernel_id == merge_kid {
debug_assert_ne!(keep_kid, merge_kid);
debug_assert!(op_map.contains_key(op_id), "merge_kernel: holder {op_id:?} not in merge kernel op list");
*kernel_id = keep_kid;
*op_id = op_map[op_id];
}
}
let store_tids: Vec<TensorId> = merge_stores.clone();
let keep_data = &mut self.kernels[keep_kid];
keep_data.outputs.extend(merge_outputs);
keep_data.loads.extend(merge_loads.iter().copied());
keep_data.stores.extend(merge_stores);
for &tid in &store_tids {
if let TensorData::PendingLeaf { depends_on, .. } = &mut self.tensors[tid] {
if *depends_on == merge_kid {
*depends_on = keep_kid;
}
}
}
#[cfg(debug_assertions)]
{
let outputs: Vec<TensorId> = self.kernels[keep_kid].outputs.iter().copied().collect();
for tid in &outputs {
match self.tensors.get(*tid) {
Some(TensorData::Eager { kernel_id, .. }) | Some(TensorData::Promoted { kernel_id, .. }) => {
debug_assert_eq!(
*kernel_id, keep_kid,
"merged output tid {tid} has kernel_id {kernel_id:?}, not keep {keep_kid:?}"
);
}
Some(t) => panic!("merge_kernel: keep kernel output tid {tid} has unexpected tensor data {t:?}"),
None => panic!("merge_kernel: keep kernel output tid {tid} was deleted from the slab (stale outputs entry)"),
}
let count = self.kernels.values().filter(|kd| kd.outputs.contains(tid)).count();
debug_assert!(count <= 1, "merged output tid {tid} is listed in {count} kernels' outputs");
}
}
Ok(op_map)
}
pub fn add_store(&mut self, x: TensorId) -> Result<(), ZyxError> {
let (kid, op_id, pending) = match self.tensors[x] {
TensorData::Eager { kernel_id, op_id, .. } => (kernel_id, op_id, KernelId::NULL),
TensorData::Promoted { kernel_id, op_id, .. } => (kernel_id, op_id, KernelId::NULL),
ref t => panic!("add_store: tensor {x} is not an eager/promoted tensor: {t:?}"),
};
debug_assert!(self.kernels[kid].outputs.contains(&x), "add_store called for tid not in outputs");
self.kernels[kid].outputs.remove(&x);
let dtype = self.dtype(x);
let add_store = self.leaf_buffer(x).is_none() && pending.is_null();
let pending = if add_store {
debug_assert!(!self.kernels[kid].loads.contains(&x), "kernel {kid:?} both loads and stores tid {x}");
let store_shape_id = self.kernels[kid].kernel.stack_shape_dims(op_id);
let dst_id =
self.kernels[kid].kernel.push_back(Op::Param { dtype, kind: ParamKind::GlobalMut, shape: store_shape_id });
self.kernels[kid].kernel.store(dst_id, op_id, OpId::NULL);
self.kernels[kid].stores.push(x);
kid
} else {
pending
};
let outputs_empty = self.kernels[kid].outputs.is_empty();
let (shape_id, rc) = match self.tensors[x] {
TensorData::Eager { shape_id, rc, .. }
| TensorData::Graph { shape_id, rc, .. }
| TensorData::Promoted { shape_id, rc, .. } => {
(shape_id, rc)
}
ref t => panic!("add_store: tensor {x} is not a kernel-backed tensor: {t:?}"),
};
self.tensors[x] = TensorData::PendingLeaf { old_buffer: None, depends_on: pending, shape_id, dtype, rc };
if outputs_empty {
self.materialize_kernel(kid)?;
}
Ok(())
}
pub fn get_or_autotune(&mut self, kernel: Kernel, buffers: &[LaunchArg]) -> Result<(DeviceProgramId, u64), ZyxError> {
let kernel_id = if let Some(&cached_kid) = self.kernel_map.get(&kernel) {
if let Some(&program_id) = self.programs.get(&cached_kid) {
let pid = ProgramId { dev: kernel.dev, program_id };
let timing = self.timings.get(&pid).copied().unwrap_or(10_000_000_000);
return Ok((program_id, timing));
}
cached_kid
} else {
let kernel_id =
KernelId::from(self.kernel_map.values().copied().max().map_or(0, |id| usize::from(id).checked_add(1).unwrap()));
let newly_inserted = self.kernel_map.insert(kernel.clone(), kernel_id).is_none();
assert!(newly_inserted);
kernel_id
};
if crate::debug_mask().sched() {
kernel.debug();
}
let device_id = kernel.dev;
#[cfg(feature = "viz")]
let sched_kernel = kernel.clone();
{
let mut n_params = 0usize;
let mut op_id = kernel.head;
while !op_id.is_null() {
if matches!(kernel.ops[op_id].op, Op::Param { .. }) {
n_params += 1;
}
op_id = kernel.next_op(op_id);
}
debug_assert_eq!(buffers.len(), n_params, "caller arg count must match kernel param count");
}
let dev_info = device_id.info();
let mut base = kernel;
base.linearize();
base.common_subexpression_elimination();
base.dead_code_elimination();
base.instruction_schedule();
{
let global_indices = base.get_group_indices();
let max_global_dims = dev_info.max_global_work_dims.len();
if global_indices.len() > max_global_dims {
let n = global_indices.len() + 1 - max_global_dims;
let indices: Vec<OpId> = global_indices.values().copied().take(n).collect();
base.merge_indices(&indices);
}
base.renumber_indices();
base.verify();
}
base.delete_zero_len_indices();
base.renumber_indices();
for _ in 0..3 {
base.default_epilogue();
}
let beam_search = crate::backend::autotune_config();
let (winner, timing) = beam_search.run_(
self,
[base],
buffers,
&Kernel::default_optimizations(),
Kernel::default_epilogue,
Kernel::base_cost,
)?;
let program_id = device_id.compile(&winner, crate::debug_mask().asm())?;
self.programs.insert(kernel_id, program_id);
self.timings.insert(ProgramId { dev: device_id, program_id }, timing);
#[cfg(feature = "viz")]
{
let kc = {
crate::viz::KernelCapture {
sched_kernel,
winner: winner.clone(),
dev_info: dev_info.clone(),
device_label: device_id.name(),
cc: match device_id {
Dev::Cuda(_) => Some(dev_info.cc),
Dev::Auto | Dev::C | Dev::Cblas | Dev::Vulkan(_) | Dev::OpenCL(_) | Dev::WGPU(_) | Dev::Dummy => None,
#[cfg(feature = "tenstorrent")]
Dev::TT(_) => None,
},
has_openmp: dev_info.has_openmp,
}
};
self.viz.record(ProgramId { dev: device_id, program_id }, kc);
}
Ok((program_id, timing))
}
fn pick_device(&self, bytes: Dim) -> Result<Dev, ZyxError> {
let mut devs: Vec<Dev> =
Dev::all().into_iter().filter(|&dev| !dev.aot_only() && dev.pool().free_bytes() >= bytes).collect();
if devs.is_empty() {
return Err(ZyxError::AllocationError(format!("no device with {bytes} bytes free").into()));
}
devs.sort_unstable_by_key(|&dev| dev.free_compute());
devs.reverse();
Ok(devs[0])
}
pub(crate) fn materialize_kernel(&mut self, kid: KernelId) -> Result<(), ZyxError> {
let mut pending_kids: Vec<KernelId> = self.kernels[kid]
.loads
.iter()
.filter_map(|&tid| match self.tensors[tid] {
TensorData::PendingLeaf { depends_on, .. } => {
debug_assert!(!depends_on.is_null(), "materialize: PendingLeaf {tid} with null depends_on");
Some(depends_on)
}
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Graph { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => None,
})
.collect();
pending_kids.sort();
pending_kids.dedup();
for pending_kid in pending_kids {
let seen: Set<TensorId> = self.kernels[pending_kid].outputs.iter().copied().collect();
for tid in seen {
self.add_store(tid)?;
}
}
for &tid in &self.kernels[kid].loads {
assert!(
matches!(self.tensors[tid], TensorData::Leaf { .. } | TensorData::GraphLeaf { .. } | TensorData::Symbolic { .. }),
"materialize: load {tid} has no buffer after pending flush: {:?}",
self.tensors[tid]
);
}
let dtypes: Map<TensorId, DType> =
self.kernels[kid].loads.iter().chain(&self.kernels[kid].stores).map(|&tid| (tid, self.dtype(tid))).collect();
for &tid in &self.kernels[kid].loads {
if let TensorData::Eager { kernel_id: k, .. } | TensorData::Promoted { kernel_id: k, .. } = &mut self.tensors[tid] {
if *k == kid {
*k = KernelId::NULL;
}
}
}
let KernelData { outputs, loads, stores, mut kernel } = unsafe { self.kernels.remove_and_return(kid) };
debug_assert!(outputs.is_empty(), "all outputs must be stored before materialize");
debug_assert!(
stores.iter().all(|&tid| matches!(self.tensors[tid], TensorData::PendingLeaf { .. })),
"materialize: not all stores are pending leafs"
);
if stores.is_empty() {
for &tid in &loads {
self.release(tid);
}
return Ok(());
}
for &tid in &loads {
assert!(
self.leaf_buffer(tid).is_some()
|| outputs.contains(&tid)
|| self.kernels.values().any(|kd| kd.outputs.contains(&tid) || kd.stores.contains(&tid))
|| self.resolve_symbolic(tid).is_some(),
"load tid {tid} not realized, not in outputs, not in any kernel; kernels loading it: {:?}",
self.kernels.iter().filter(|(_, kd)| kd.loads.contains(&tid)).map(|(k, _)| k).collect::<Vec<_>>(),
);
}
#[cfg(debug_assertions)]
{
for &tid in &stores {
let count = self.kernels.values().filter(|kd| kd.outputs.contains(&tid)).count();
debug_assert!(count <= 1, "store tid={tid} is in {count} kernels' outputs");
}
let mut listed_in_multiple = Vec::new();
let mut counted: Map<TensorId, usize> = Map::with_hasher(BuildHasherDefault::new());
for kd in self.kernels.values() {
for &tid in &kd.outputs {
*counted.entry(tid).or_insert(0) += 1;
}
}
for (&tid, &count) in &counted {
if count > 1 {
listed_in_multiple.push(tid);
}
}
debug_assert!(
listed_in_multiple.is_empty(),
"inventory desync: tensors listed in multiple kernels' outputs: {listed_in_multiple:?}"
);
}
for &load in &loads {
if self.leaf_buffer(load).is_some() || self.resolve_symbolic(load).is_some() {
continue;
}
let pending = match self.tensors[load] {
TensorData::PendingLeaf { depends_on, .. } => depends_on,
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Graph { .. }
| TensorData::Promoted { .. }
| TensorData::Symbolic { .. } => KernelId::NULL,
};
if pending.is_null() {
continue;
}
let outputs: Set<TensorId> = self.kernels[pending].outputs.iter().copied().collect();
if outputs.is_empty() {
self.materialize_kernel(pending)?;
continue;
}
for output in outputs {
self.add_store(output)?;
}
}
for &load in &loads {
if self.leaf_buffer(load).is_some() || self.resolve_symbolic(load).is_some() {
continue;
}
if matches!(
self.tensors[load],
TensorData::Eager { .. } | TensorData::Promoted { .. } | TensorData::PendingLeaf { .. }
) {
self.add_store(load)?;
}
}
debug_assert!(
loads.iter().all(|&tid| self.leaf_buffer(tid).is_some() || self.resolve_symbolic(tid).is_some()),
"all loads must be realized after recursive materialization"
);
let mut store_pools: BTreeSet<Pool> = BTreeSet::new();
for &tid in &stores {
if let TensorData::PendingLeaf { old_buffer: Some(buf_id), .. } = self.tensors[tid] {
store_pools.insert(buf_id.pool);
}
}
let out_bytes: Dim = stores
.iter()
.filter(|&&tid| matches!(self.tensors[tid], TensorData::PendingLeaf { old_buffer: None, .. }))
.map(|&tid| {
let dtype = dtypes[&tid];
(self.resolve_shape(tid).iter().product::<Dim>() * dtype.bit_size() as Dim + 7) / 8
})
.sum();
let (dev_id, pool_id) = if kernel.dev != Dev::Auto {
let dev_id = kernel.dev;
if store_pools.len() > 1 || !store_pools.iter().all(|&pool| pool == dev_id.pool()) {
return Err(ZyxError::AllocationError(
format!(
"stores {store_pools:?} do not match the kernel's pinned device {dev_id:?} and cannot span multiple pools",
)
.into(),
));
}
(dev_id, dev_id.pool())
} else if store_pools.len() == 1 {
let pool_id = *store_pools.iter().next().unwrap();
let dev_id = Dev::all().into_iter().find(|dev| dev.pool() == pool_id);
match dev_id {
Some(dev_id) => (dev_id, pool_id),
None => {
let dev_id = self.pick_device(out_bytes)?;
(dev_id, dev_id.pool())
}
}
} else if store_pools.is_empty() {
let mut loaded_bytes: Map<Pool, Dim> = Map::default();
for &tid in &loads {
if let Some(buf_id) = self.leaf_buffer(tid) {
let dtype = dtypes[&tid];
*loaded_bytes.entry(buf_id.pool).or_insert(0) +=
(self.resolve_shape(tid).iter().product::<Dim>() * dtype.bit_size() as Dim + 7) / 8;
}
}
let mut best: Option<(Dev, Pool)> = None;
for dev_id in Dev::all() {
if dev_id.aot_only() {
continue;
}
let pool_id = dev_id.pool();
if pool_id.free_bytes() < out_bytes {
continue;
}
let bytes = loaded_bytes.get(&pool_id).copied().unwrap_or(0);
if best.is_none_or(|(_, best_pool)| loaded_bytes.get(&best_pool).copied().unwrap_or(0) < bytes) {
best = Some((dev_id, pool_id));
}
}
match best {
Some((dev_id, pool_id)) => (dev_id, pool_id),
None => {
let dev_id = self.pick_device(out_bytes)?;
(dev_id, dev_id.pool())
}
}
} else {
return Err(ZyxError::AllocationError(
format!("stores span multiple pools {store_pools:?}; a kernel can only touch memory of a single pool").into(),
));
};
kernel.dev = dev_id;
kernel.dev_info = Some(dev_id.info());
for &tid in &loads {
let Some(buf_id) = self.leaf_buffer(tid) else { continue };
if buf_id.pool != pool_id {
let src = buf_id.buffer_id;
let bytes =
(self.resolve_shape(tid).iter().product::<Dim>() as usize * dtypes[&tid].bit_size() as usize).div_ceil(8);
let alloc_bytes = bytes + dtypes[&tid].bit_size() as usize / 8;
let dst = pool_id.allocate(alloc_bytes as Dim)?;
let dst_global = Buffer { pool: pool_id, buffer_id: dst };
debug_assert_ne!(buf_id.pool, pool_id, "pool_to_pool across the same pool is disallowed");
pool_id.pool_to_pool(buf_id.pool, src, dst)?;
buf_id.pool.release(src);
match &mut self.tensors[tid] {
TensorData::Leaf { buffer: buffer_id, .. } | TensorData::GraphLeaf { buffer: buffer_id, .. } => {
*buffer_id = dst_global;
}
ref t => panic!("materialize: moved load {tid} has no buffer field: {t:?}"),
}
}
}
for &tid in &stores {
let Some(buf_id) = (match self.tensors[tid] {
TensorData::PendingLeaf { old_buffer, .. } => old_buffer,
ref t => panic!("materialize: store {tid} is not a pending leaf: {t:?}"),
}) else {
continue;
};
if buf_id.pool != pool_id {
let src = buf_id.buffer_id;
let bytes =
(self.resolve_shape(tid).iter().product::<Dim>() as usize * dtypes[&tid].bit_size() as usize).div_ceil(8);
let alloc_bytes = bytes as Dim + Dim::from(dtypes[&tid].bit_size() / 8);
let dst = pool_id.allocate(alloc_bytes)?;
let dst_global = Buffer { pool: pool_id, buffer_id: dst };
debug_assert_ne!(buf_id.pool, pool_id, "pool_to_pool across the same pool is disallowed");
pool_id.pool_to_pool(buf_id.pool, src, dst)?;
buf_id.pool.release(src);
match &mut self.tensors[tid] {
TensorData::PendingLeaf { old_buffer: Some(buffer_id), .. } => {
*buffer_id = dst_global;
}
ref t => panic!("materialize: moved store {tid} is not a pending leaf: {t:?}"),
}
}
}
let mut kernel_buffers = BTreeSet::new();
for &tid in &loads {
if matches!(self.tensors[tid], TensorData::Symbolic { expr, .. } if matches!(self.exprs[expr], Expr::Variable { .. }))
{
continue;
}
kernel_buffers.insert(self.leaf_buffer(tid).expect("materialize: load without buffer after pending flush"));
}
for &tid in &stores {
if let TensorData::PendingLeaf { old_buffer: Some(buf), .. } = self.tensors[tid] {
kernel_buffers.insert(buf);
continue;
}
let bytes = (self.resolve_shape(tid).iter().product::<Dim>() as usize * dtypes[&tid].bit_size() as usize).div_ceil(8);
let alloc_bytes = bytes as Dim + Dim::from(dtypes[&tid].bit_size() / 8);
let buf = pool_id.allocate(alloc_bytes)?;
let global_id = Buffer { pool: pool_id, buffer_id: buf };
kernel_buffers.insert(global_id);
match &mut self.tensors[tid] {
TensorData::PendingLeaf { old_buffer: slot @ None, .. } => {
*slot = Some(global_id);
}
ref t => panic!("materialize: store {tid} is not a pending leaf: {t:?}"),
}
}
for &tid in &stores {
debug_assert!(
matches!(self.tensors[tid], TensorData::PendingLeaf { old_buffer: Some(_), .. }),
"materialize: store tid {tid} has no buffer after realization"
);
}
#[cfg(debug_assertions)]
{
let (mut n_non_mut, mut n_mut) = (0usize, 0usize);
let mut p = kernel.head;
while !p.is_null() {
if let Op::Param { kind, .. } = &kernel.ops[p].op {
match kind {
ParamKind::GlobalMut => n_mut += 1,
ParamKind::Global | ParamKind::Variable => n_non_mut += 1,
}
}
p = kernel.next_op(p);
}
assert_eq!(n_non_mut, loads.len(), "materialize: {} non-store defines but {} load entries", n_non_mut, loads.len());
assert!(n_mut <= stores.len(), "materialize: {} GlobalMut defines but only {} stores", n_mut, stores.len());
}
let mut buffers: Vec<LaunchArg> = Vec::new();
for &tid in &loads {
let var_value = match self.tensors[tid] {
TensorData::Symbolic { expr, .. } => match self.exprs[expr] {
Expr::Variable { value } => Some(value),
ref e => panic!("materialize: symbolic load {tid} is not a variable: {e:?}"),
},
TensorData::Eager { .. }
| TensorData::Leaf { .. }
| TensorData::PendingLeaf { .. }
| TensorData::GraphLeaf { .. }
| TensorData::Graph { .. }
| TensorData::Promoted { .. } => None,
};
if let Some(value) = var_value {
buffers.push(LaunchArg::Variable(value));
} else {
buffers.push(LaunchArg::Buffer(self.leaf_buffer(tid).expect("materialize: load without buffer").buffer_id));
}
}
for &tid in &stores {
let buf = match self.tensors[tid] {
TensorData::PendingLeaf { old_buffer: Some(buf), .. } => buf,
ref t => panic!("materialize: store {tid} has no buffer after realization: {t:?}"),
};
buffers.push(LaunchArg::Buffer(buf.buffer_id));
}
let (dev_prog, _timing) = self.get_or_autotune(kernel, &buffers)?;
dev_id.launch(dev_prog, &buffers)?;
for &tid in &stores {
let (shape_id, dtype, rc, buf) = match self.tensors[tid] {
TensorData::PendingLeaf { shape_id, dtype, rc, old_buffer: Some(buf), .. } => (shape_id, dtype, rc, buf),
ref t => panic!("materialize: store {tid} is not a realized pending leaf: {t:?}"),
};
self.tensors[tid] = TensorData::Leaf { shape_id, dtype, buffer: buf, rc };
}
for &tid in &loads {
self.release(tid);
}
Ok(())
}
#[cfg(test)]
#[allow(unused)]
pub fn live_inventory(&self) -> usize {
self.tensors.iter().count()
}
}