use crate::arena::{ContextStack, DataValue};
use crate::{CompiledNode, Engine, Error, Result};
use bumpalo::Bump;
use datavalue::{DType, DataTensor, TensorError};
mod construct;
mod read;
mod shape_ops;
pub(crate) use construct::{
evaluate_full, evaluate_one_hot, evaluate_rle_expand, evaluate_scatter, evaluate_tensor,
evaluate_zeros,
};
pub(crate) use read::{
evaluate_argmax, evaluate_cast, evaluate_dtype, evaluate_normalize, evaluate_shape,
evaluate_to_list,
};
pub(crate) use shape_ops::{
evaluate_concat, evaluate_crop, evaluate_gather, evaluate_pad, evaluate_reshape,
evaluate_stack, evaluate_transpose, evaluate_unstack,
};
#[inline]
fn charge(ctx: &mut ContextStack<'_>, elements: u64) -> Result<()> {
ctx.charge(elements)
}
#[inline]
fn cost(a: usize, b: usize) -> u64 {
a.max(b) as u64
}
#[inline]
fn wrap(e: TensorError) -> Error {
Error::wrap(e)
}
#[inline]
fn bad(msg: &'static str) -> Error {
Error::invalid_arguments(msg)
}
fn element_error(index: usize, expected: DType) -> Error {
wrap(TensorError::Element { index, expected })
}
#[inline]
fn arg<'a>(
args: &'a [CompiledNode],
i: usize,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
let node = args.get(i).ok_or_else(|| bad("missing argument"))?;
engine.dispatch_node(node, ctx, arena)
}
#[inline]
fn opt_arg<'a>(
args: &'a [CompiledNode],
i: usize,
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<Option<&'a DataValue<'a>>> {
match args.get(i) {
None => Ok(None),
Some(node) => engine.dispatch_node(node, ctx, arena).map(Some),
}
}
#[inline]
fn at_most(args: &[CompiledNode], n: usize) -> Result<()> {
if args.len() > n {
return Err(bad("too many arguments"));
}
Ok(())
}
fn as_tensor<'a>(v: &DataValue<'a>, arena: &'a Bump) -> Result<DataTensor<'a>> {
match v {
DataValue::Tensor(t) => Ok(**t),
DataValue::Object(_) => DataTensor::from_json_value_in(v, arena).map_err(wrap),
_ => Err(bad("expected a tensor")),
}
}
fn as_dtype(v: &DataValue<'_>) -> Result<DType> {
let DataValue::String(s) = v else {
return Err(bad("dtype must be a string"));
};
DType::from_name(s).ok_or_else(|| wrap(TensorError::UnknownDType((*s).to_string())))
}
fn as_shape<'a>(v: &DataValue<'_>, arena: &'a Bump) -> Result<&'a [usize]> {
match v {
DataValue::Number(_) => Ok(arena.alloc_slice_copy(&[as_usize(v)?])),
DataValue::Array(items) => {
let mut out = crate::arena::bvec::<usize>(arena, items.len());
for it in *items {
out.push(as_usize(it)?);
}
Ok(out.into_bump_slice())
}
_ => Err(bad("shape must be an array of non-negative integers")),
}
}
fn as_i64_list<'a>(v: &DataValue<'_>, arena: &'a Bump) -> Result<&'a [i64]> {
let DataValue::Array(items) = v else {
return Err(bad("expected an array of integers"));
};
let mut out = crate::arena::bvec::<i64>(arena, items.len());
for it in *items {
out.push(as_i64(it)?);
}
Ok(out.into_bump_slice())
}
fn as_usize(v: &DataValue<'_>) -> Result<usize> {
let n = as_i64(v)?;
usize::try_from(n).map_err(|_| bad("expected a non-negative integer"))
}
fn as_i64(v: &DataValue<'_>) -> Result<i64> {
match v {
DataValue::Number(n) => {
let f = n.as_f64();
if f.fract() != 0.0 || !f.is_finite() {
return Err(bad("expected a whole number"));
}
n.as_i64()
.or_else(|| {
(f >= -(2f64.powi(63)) && f < 2f64.powi(63)).then_some(f as i64)
})
.ok_or_else(|| bad("integer out of range"))
}
_ => Err(bad("expected a whole number")),
}
}
fn as_f64(v: &DataValue<'_>) -> Result<f64> {
match v {
DataValue::Number(n) => Ok(n.as_f64()),
_ => Err(bad("expected a number")),
}
}
#[inline]
fn resolve_index(i: i64, extent: usize) -> Option<usize> {
let resolved = if i < 0 { i + extent as i64 } else { i };
usize::try_from(resolved).ok().filter(|r| *r < extent)
}
fn as_axis(v: &DataValue<'_>, rank: usize, extra: usize) -> Result<usize> {
resolve_index(as_i64(v)?, (rank + extra).max(1)).ok_or_else(|| bad("axis out of range"))
}
fn numel_of(shape: &[usize]) -> Result<usize> {
shape
.iter()
.try_fold(1usize, |acc, d| acc.checked_mul(*d))
.ok_or_else(|| wrap(TensorError::ShapeOverflow))
}
fn strides_of<'a>(shape: &[usize], arena: &'a Bump) -> &'a [usize] {
let mut s = crate::arena::bvec::<usize>(arena, shape.len());
s.resize(shape.len(), 1);
for i in (0..shape.len().saturating_sub(1)).rev() {
s[i] = s[i + 1] * shape[i + 1];
}
s.into_bump_slice()
}
#[inline]
fn advance(idx: &mut [usize], shape: &[usize]) -> bool {
for ax in (0..shape.len()).rev() {
idx[ax] += 1;
if idx[ax] < shape[ax] {
return true;
}
idx[ax] = 0;
}
false
}
#[inline]
fn split_axis(shape: &[usize], axis: usize) -> (usize, usize, usize) {
(
shape[..axis].iter().product(),
shape[axis],
shape[axis + 1..].iter().product(),
)
}
fn shape_without_axis<'a>(shape: &[usize], axis: usize, arena: &'a Bump) -> &'a [usize] {
let mut out = crate::arena::bvec::<usize>(arena, shape.len() - 1);
out.extend_from_slice(&shape[..axis]);
out.extend_from_slice(&shape[axis + 1..]);
out.into_bump_slice()
}
#[inline]
fn finish<'a>(t: DataTensor<'a>, arena: &'a Bump) -> Result<&'a DataValue<'a>> {
Ok(arena.alloc(DataValue::tensor_in(t, arena)))
}
#[inline]
fn finish_bytes<'a>(
dtype: DType,
shape: &'a [usize],
bytes: &'a [u8],
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
finish(
DataTensor::from_bytes(dtype, shape, bytes).map_err(wrap)?,
arena,
)
}
#[inline]
fn finish_slice<'a, T: datavalue::Element>(
shape: &'a [usize],
data: bumpalo::collections::Vec<'a, T>,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
finish(
DataTensor::from_slice(shape, data.into_bump_slice()).map_err(wrap)?,
arena,
)
}
pub(crate) trait Scalar: datavalue::Element + PartialOrd {
const ZERO: Self;
const ONE: Self;
fn from_value(v: &DataValue<'_>) -> Option<Self>;
fn from_f64_saturating(v: f64) -> Self;
fn to_f64(self) -> f64;
}
macro_rules! int_scalar {
($($t:ty),* $(,)?) => {$(
impl Scalar for $t {
const ZERO: Self = 0;
const ONE: Self = 1;
fn from_value(v: &DataValue<'_>) -> Option<Self> {
let DataValue::Number(n) = v else { return None };
if let Some(i) = n.as_i64() {
return Self::try_from(i).ok();
}
let f = n.as_f64();
if f.fract() != 0.0 || !f.is_finite() {
return None;
}
let as_int = f as i128;
(as_int as f64 == f).then(|| Self::try_from(as_int).ok())?
}
fn from_f64_saturating(v: f64) -> Self {
v as Self
}
fn to_f64(self) -> f64 {
self as f64
}
}
)*};
}
int_scalar!(i8, u8, i16, u16, i32, u32, i64, u64);
impl Scalar for f32 {
const ZERO: Self = 0.0;
const ONE: Self = 1.0;
fn from_value(v: &DataValue<'_>) -> Option<Self> {
let DataValue::Number(n) = v else { return None };
let f = n.as_f64();
let narrowed = f as f32;
(narrowed.is_finite() || !f.is_finite()).then_some(narrowed)
}
fn from_f64_saturating(v: f64) -> Self {
v as f32
}
fn to_f64(self) -> f64 {
self as f64
}
}
impl Scalar for f64 {
const ZERO: Self = 0.0;
const ONE: Self = 1.0;
fn from_value(v: &DataValue<'_>) -> Option<Self> {
match v {
DataValue::Number(n) => Some(n.as_f64()),
_ => None,
}
}
fn from_f64_saturating(v: f64) -> Self {
v
}
fn to_f64(self) -> f64 {
self
}
}
impl Scalar for bool {
const ZERO: Self = false;
const ONE: Self = true;
fn from_value(v: &DataValue<'_>) -> Option<Self> {
match v {
DataValue::Bool(b) => Some(*b),
DataValue::Number(n) => {
let f = n.as_f64();
if f == 0.0 {
Some(false)
} else if f == 1.0 {
Some(true)
} else {
None
}
}
_ => None,
}
}
fn from_f64_saturating(v: f64) -> Self {
v != 0.0 && !v.is_nan()
}
fn to_f64(self) -> f64 {
if self { 1.0 } else { 0.0 }
}
}
#[cfg(feature = "tensor-half")]
macro_rules! half_scalar {
($($t:ty),* $(,)?) => {$(
impl Scalar for $t {
const ZERO: Self = <$t>::from_f32_const(0.0);
const ONE: Self = <$t>::from_f32_const(1.0);
fn from_value(v: &DataValue<'_>) -> Option<Self> {
let DataValue::Number(n) = v else { return None };
let f = n.as_f64();
let narrowed = <$t>::from_f64(f);
(narrowed.is_finite() || !f.is_finite()).then_some(narrowed)
}
fn from_f64_saturating(v: f64) -> Self {
<$t>::from_f64(v)
}
fn to_f64(self) -> f64 {
<$t>::to_f64(self)
}
}
)*};
}
#[cfg(feature = "tensor-half")]
half_scalar!(datavalue::half::f16, datavalue::half::bf16);
macro_rules! by_dtype {
($dtype:expr, $f:ident $(, $arg:expr)* $(,)?) => {{
match $dtype {
::datavalue::DType::Bool => $f::<bool>($($arg),*),
::datavalue::DType::I8 => $f::<i8>($($arg),*),
::datavalue::DType::U8 => $f::<u8>($($arg),*),
::datavalue::DType::I16 => $f::<i16>($($arg),*),
::datavalue::DType::U16 => $f::<u16>($($arg),*),
::datavalue::DType::I32 => $f::<i32>($($arg),*),
::datavalue::DType::U32 => $f::<u32>($($arg),*),
::datavalue::DType::I64 => $f::<i64>($($arg),*),
::datavalue::DType::U64 => $f::<u64>($($arg),*),
::datavalue::DType::F32 => $f::<f32>($($arg),*),
::datavalue::DType::F64 => $f::<f64>($($arg),*),
#[cfg(feature = "tensor-half")]
::datavalue::DType::F16 => $f::<::datavalue::half::f16>($($arg),*),
#[cfg(feature = "tensor-half")]
::datavalue::DType::BF16 => $f::<::datavalue::half::bf16>($($arg),*),
other => ::core::result::Result::Err($crate::operators::tensor::wrap(
::datavalue::TensorError::UnsupportedDType(other),
)),
}
}};
}
pub(crate) use by_dtype;