use super::{
Scalar, arg, as_axis, as_dtype, as_f64, as_tensor, at_most, bad, by_dtype, charge, finish,
finish_slice, opt_arg, shape_without_axis, split_axis, wrap,
};
use crate::arena::{ContextStack, DataValue, bvec};
use crate::{CompiledNode, Engine, Result};
use bumpalo::Bump;
use datavalue::{DType, DataTensor, NumberValue};
#[inline]
fn elements<'t, T: Scalar>(t: &'t DataTensor<'_>) -> Result<&'t [T]> {
t.as_slice::<T>()
.ok_or_else(|| bad("tensor: dtype does not match its payload"))
}
fn widen<'a>(t: DataTensor<'_>, arena: &'a Bump) -> Result<&'a [f64]> {
by_dtype!(t.dtype(), widen_impl, t, arena)
}
fn widen_impl<'a, T: Scalar>(t: DataTensor<'_>, arena: &'a Bump) -> Result<&'a [f64]> {
let src = elements::<T>(&t)?;
let mut out = bvec::<f64>(arena, src.len());
out.extend(src.iter().map(|v| v.to_f64()));
Ok(out.into_bump_slice())
}
fn narrow<'a>(
dtype: DType,
values: &[f64],
shape: &'a [usize],
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
by_dtype!(dtype, narrow_impl, values, shape, arena)
}
fn narrow_impl<'a, T: Scalar>(
values: &[f64],
shape: &'a [usize],
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
let mut out = bvec::<T>(arena, values.len());
out.extend(values.iter().map(|v| T::from_f64_saturating(*v)));
finish_slice(shape, out, arena)
}
pub(crate) fn evaluate_cast<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 2)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
let dtype = as_dtype(arg(args, 1, ctx, engine, arena)?)?;
charge(ctx, t.numel() as u64)?;
if dtype == t.dtype() {
return finish(t, arena);
}
let values = widen(t, arena)?;
narrow(dtype, values, t.shape(), arena)
}
pub(crate) fn evaluate_normalize<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 3)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
let mean = as_f64(arg(args, 1, ctx, engine, arena)?)?;
let scale = match opt_arg(args, 2, ctx, engine, arena)? {
Some(v) => as_f64(v)?,
None => 1.0,
};
charge(ctx, t.numel() as u64)?;
by_dtype!(t.dtype(), normalize_impl, t, mean, scale, arena)
}
fn normalize_impl<'a, T: Scalar>(
t: DataTensor<'a>,
mean: f64,
scale: f64,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
let src = elements::<T>(&t)?;
let mut out = bvec::<f32>(arena, src.len());
out.extend(src.iter().map(|v| ((v.to_f64() - mean) * scale) as f32));
finish_slice(t.shape(), out, arena)
}
pub(crate) fn evaluate_argmax<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 2)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
if t.ndim() == 0 {
return Err(bad("argmax: a 0-d tensor has no axis to reduce"));
}
let axis = as_axis(arg(args, 1, ctx, engine, arena)?, t.ndim(), 0)?;
if t.shape()[axis] == 0 {
return Err(bad("argmax: cannot reduce a zero-length axis"));
}
charge(ctx, t.numel() as u64)?;
by_dtype!(t.dtype(), argmax_impl, t, axis, arena)
}
fn argmax_impl<'a, T: Scalar>(
t: DataTensor<'a>,
axis: usize,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
let src = elements::<T>(&t)?;
let (outer, extent, inner) = split_axis(t.shape(), axis);
let mut out = bvec::<i64>(arena, outer * inner);
for o in 0..outer {
for i in 0..inner {
let at = |k: usize| src[((o * extent) + k) * inner + i];
let mut best = 0usize;
let mut best_v = at(0);
for k in 1..extent {
if at(k) > best_v {
best_v = at(k);
best = k;
}
}
out.push(best as i64);
}
}
let shape = shape_without_axis(t.shape(), axis, arena);
let reduced = DataTensor::from_slice(shape, out.into_bump_slice()).map_err(wrap)?;
Ok(arena.alloc(reduced.to_nested_in(arena).map_err(wrap)?))
}
pub(crate) fn evaluate_to_list<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 1)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
charge(ctx, t.numel() as u64)?;
Ok(arena.alloc(t.to_nested_in(arena).map_err(wrap)?))
}
pub(crate) fn evaluate_shape<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 1)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
charge(ctx, 1)?;
let dims = arena.alloc_slice_fill_iter(
t.shape()
.iter()
.map(|d| DataValue::Number(NumberValue::Integer(*d as i64))),
);
Ok(arena.alloc(DataValue::Array(dims)))
}
pub(crate) fn evaluate_dtype<'a>(
args: &'a [CompiledNode],
ctx: &mut ContextStack<'a>,
engine: &Engine,
arena: &'a Bump,
) -> Result<&'a DataValue<'a>> {
at_most(args, 1)?;
let t = as_tensor(arg(args, 0, ctx, engine, arena)?, arena)?;
charge(ctx, 1)?;
Ok(arena.alloc(DataValue::String(t.dtype().name())))
}