use num_complex::Complex64;
use tenferro_ops::dim_expr::DimExpr;
use tenferro_ops::ShapeExtent;
use tenferro_runtime::error::Error;
use tenferro_runtime::extension::{ExecInstruction, ExecOp, ExecProgram};
use tenferro_runtime::{GraphCompiler, GraphExecutor, TensorRead, TracedTensor};
use tenferro_tensor::{
BackendCachedDot, BackendRuntimeCache, BackendSessionHost, CompareDir, DType, DotGeneralConfig,
GatherConfig, PadConfig, ScatterConfig, SliceConfig, Tensor, TensorAnalytic, TensorBackend,
TensorBuffer, TensorDeviceTransfer, TensorDot, TensorElementwise, TensorFusion, TensorIndexing,
TensorReduction, TensorStructural, TypedTensor,
};
fn dim_shape(shape: &[usize]) -> Vec<DimExpr> {
DimExpr::from_concrete(shape)
}
fn empty_extents(output_count: usize) -> Vec<Vec<ShapeExtent<DimExpr>>> {
vec![Vec::new(); output_count]
}
fn scalar_tensor(value: f64) -> Tensor {
Tensor::F64(TypedTensor::from_vec_col_major(vec![], vec![value]).unwrap())
}
fn f64_tensor(shape: Vec<usize>, data: Vec<f64>) -> Tensor {
Tensor::F64(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn scalar_value(tensor: &Tensor) -> f64 {
match tensor {
Tensor::F64(inner) => inner.host_data().unwrap()[0],
other => panic!("expected scalar f64 tensor, got {other:?}"),
}
}
fn scalar_c64_value(tensor: &Tensor) -> Complex64 {
match tensor {
Tensor::C64(inner) => inner.host_data().unwrap()[0],
other => panic!("expected scalar c64 tensor, got {other:?}"),
}
}
fn gather_config() -> GatherConfig {
GatherConfig {
offset_dims: vec![],
collapsed_slice_dims: vec![0],
start_index_map: vec![0],
index_vector_dim: 1,
slice_sizes: vec![1],
}
}
fn scatter_config() -> ScatterConfig {
ScatterConfig {
update_window_dims: vec![],
inserted_window_dims: vec![0],
scatter_dims_to_operand_dims: vec![0],
index_vector_dim: 1,
}
}
fn pad_config() -> PadConfig {
PadConfig {
edge_padding_low: vec![1],
edge_padding_high: vec![1],
interior_padding: vec![0],
}
}
fn single_instruction_program(op: ExecOp, input_count: usize) -> ExecProgram {
ExecProgram {
instructions: vec![ExecInstruction {
op,
input_slots: (0..input_count).collect(),
output_slots: vec![input_count],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![false; input_count],
}],
input_slots: (0..input_count).collect(),
output_slots: vec![input_count],
n_slots: input_count + 1,
}
}
#[derive(Default)]
struct FakeTensorBackend {
calls: Vec<&'static str>,
error_on: Option<&'static str>,
reclaimed: usize,
last_gather_slice_sizes: Option<Vec<usize>>,
}
impl FakeTensorBackend {
fn result(&mut self, name: &'static str, value: f64) -> tenferro_tensor::Result<Tensor> {
self.calls.push(name);
if self.error_on == Some(name) {
return Err(tenferro_tensor::Error::backend_failure(
name,
"injected failure",
));
}
Ok(scalar_tensor(value))
}
}
impl BackendRuntimeCache for FakeTensorBackend {
type RuntimeCache = ();
}
impl TensorElementwise for FakeTensorBackend {
fn add(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("add", 1.0)
}
fn mul(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("mul", 2.0)
}
fn neg(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("neg", 3.0)
}
fn conj(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("conj", 4.0)
}
fn div(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("div", 5.0)
}
fn abs(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("abs", 6.0)
}
fn sign(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("sign", 7.0)
}
fn maximum(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("maximum", 8.0)
}
fn minimum(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("minimum", 9.0)
}
fn compare(
&mut self,
_lhs: &Tensor,
_rhs: &Tensor,
_dir: &CompareDir,
) -> tenferro_tensor::Result<Tensor> {
self.result("compare", 10.0)
}
fn select(
&mut self,
_pred: &Tensor,
_on_true: &Tensor,
_on_false: &Tensor,
) -> tenferro_tensor::Result<Tensor> {
self.result("select", 11.0)
}
fn clamp(
&mut self,
_input: &Tensor,
_lower: &Tensor,
_upper: &Tensor,
) -> tenferro_tensor::Result<Tensor> {
self.result("clamp", 12.0)
}
}
impl TensorAnalytic for FakeTensorBackend {
fn exp(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("exp", 13.0)
}
fn log(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("log", 14.0)
}
fn sin(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("sin", 15.0)
}
fn cos(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("cos", 16.0)
}
fn tanh(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("tanh", 17.0)
}
fn sqrt(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("sqrt", 18.0)
}
fn rsqrt(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("rsqrt", 19.0)
}
fn pow(&mut self, _lhs: &Tensor, _rhs: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("pow", 20.0)
}
fn expm1(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("expm1", 21.0)
}
fn log1p(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.result("log1p", 22.0)
}
}
impl TensorStructural for FakeTensorBackend {
fn transpose(&mut self, _input: &Tensor, _perm: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("transpose", 23.0)
}
fn reshape(&mut self, _input: &Tensor, _shape: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reshape", 24.0)
}
fn broadcast_in_dim(
&mut self,
_input: &Tensor,
_shape: &[usize],
_dims: &[usize],
) -> tenferro_tensor::Result<Tensor> {
self.result("broadcast_in_dim", 25.0)
}
fn cast(&mut self, _input: &Tensor, _to: DType) -> tenferro_tensor::Result<Tensor> {
self.result("cast", 25.5)
}
fn extract_diagonal(
&mut self,
_input: &Tensor,
_axis_a: usize,
_axis_b: usize,
) -> tenferro_tensor::Result<Tensor> {
self.result("extract_diagonal", 26.0)
}
fn embed_diagonal(
&mut self,
_input: &Tensor,
_axis_a: usize,
_axis_b: usize,
) -> tenferro_tensor::Result<Tensor> {
self.result("embed_diagonal", 27.0)
}
fn tril(&mut self, _input: &Tensor, _k: i64) -> tenferro_tensor::Result<Tensor> {
self.result("tril", 27.5)
}
fn triu(&mut self, _input: &Tensor, _k: i64) -> tenferro_tensor::Result<Tensor> {
self.result("triu", 27.75)
}
}
impl TensorReduction for FakeTensorBackend {
fn reduce_sum(&mut self, _input: &Tensor, _axes: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reduce_sum", 28.0)
}
fn reduce_prod(&mut self, _input: &Tensor, _axes: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reduce_prod", 29.0)
}
fn reduce_max(&mut self, _input: &Tensor, _axes: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reduce_max", 30.0)
}
fn reduce_min(&mut self, _input: &Tensor, _axes: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reduce_min", 31.0)
}
}
impl TensorDot for FakeTensorBackend {
fn dot_general(
&mut self,
_lhs: &Tensor,
_rhs: &Tensor,
_config: &DotGeneralConfig,
) -> tenferro_tensor::Result<Tensor> {
self.result("dot_general", 32.0)
}
}
impl TensorIndexing for FakeTensorBackend {
fn gather(
&mut self,
_operand: &Tensor,
_start_indices: &Tensor,
config: &GatherConfig,
) -> tenferro_tensor::Result<Tensor> {
self.last_gather_slice_sizes = Some(config.slice_sizes.clone());
self.result("gather", 33.0)
}
fn scatter(
&mut self,
_operand: &Tensor,
_scatter_indices: &Tensor,
_updates: &Tensor,
_config: &ScatterConfig,
) -> tenferro_tensor::Result<Tensor> {
self.result("scatter", 34.0)
}
fn slice(&mut self, _input: &Tensor, _config: &SliceConfig) -> tenferro_tensor::Result<Tensor> {
self.result("slice", 35.0)
}
fn dynamic_slice(
&mut self,
_input: &Tensor,
_starts: &Tensor,
_slice_sizes: &[usize],
) -> tenferro_tensor::Result<Tensor> {
self.result("dynamic_slice", 36.0)
}
fn dynamic_update_slice(
&mut self,
_operand: &Tensor,
_update: &Tensor,
_starts: &Tensor,
) -> tenferro_tensor::Result<Tensor> {
self.result("dynamic_update_slice", 36.5)
}
fn pad(&mut self, _input: &Tensor, _config: &PadConfig) -> tenferro_tensor::Result<Tensor> {
self.result("pad", 37.0)
}
fn concatenate(
&mut self,
_inputs: &[&Tensor],
_axis: usize,
) -> tenferro_tensor::Result<Tensor> {
self.result("concatenate", 38.0)
}
fn reverse(&mut self, _input: &Tensor, _axes: &[usize]) -> tenferro_tensor::Result<Tensor> {
self.result("reverse", 39.0)
}
}
impl TensorBuffer for FakeTensorBackend {
fn reclaim_buffer(&mut self, _tensor: Tensor) {
self.reclaimed += 1;
}
}
impl TensorFusion for FakeTensorBackend {}
impl TensorDeviceTransfer for FakeTensorBackend {}
impl BackendCachedDot for FakeTensorBackend {}
impl BackendSessionHost for FakeTensorBackend {}
impl TensorBackend for FakeTensorBackend {}
#[test]
fn eval_exec_ir_dispatches_tensor_ops_to_backend_methods() {
let cases = vec![
(ExecOp::Transpose { perm: vec![0] }, 1, "transpose", 23.0),
(
ExecOp::Reshape {
shape: dim_shape(&[1]),
},
1,
"reshape",
24.0,
),
(
ExecOp::BroadcastInDim {
shape: dim_shape(&[1]),
dims: vec![0],
},
1,
"broadcast_in_dim",
25.0,
),
(ExecOp::Convert { to: DType::C64 }, 1, "cast", 25.5),
(
ExecOp::DotGeneral(DotGeneralConfig {
lhs_contracting_dims: vec![0],
rhs_contracting_dims: vec![0],
lhs_batch_dims: vec![],
rhs_batch_dims: vec![],
}),
2,
"dot_general",
32.0,
),
(ExecOp::ReduceSum { axes: vec![0] }, 1, "reduce_sum", 28.0),
(
ExecOp::ExtractDiag {
axis_a: 0,
axis_b: 1,
},
1,
"extract_diagonal",
26.0,
),
(
ExecOp::EmbedDiag {
axis_a: 0,
axis_b: 1,
},
1,
"embed_diagonal",
27.0,
),
(ExecOp::Tril { k: -1 }, 1, "tril", 27.5),
(ExecOp::Triu { k: 1 }, 1, "triu", 27.75),
(ExecOp::Add, 2, "add", 1.0),
(ExecOp::Multiply, 2, "mul", 2.0),
(ExecOp::Negate, 1, "neg", 3.0),
(ExecOp::Conj, 1, "conj", 4.0),
(ExecOp::Divide, 2, "div", 5.0),
(ExecOp::Abs, 1, "abs", 6.0),
(ExecOp::Sign, 1, "sign", 7.0),
(ExecOp::Maximum, 2, "maximum", 8.0),
(ExecOp::Minimum, 2, "minimum", 9.0),
(ExecOp::Compare(CompareDir::Eq), 2, "compare", 10.0),
(ExecOp::Select, 3, "select", 11.0),
(ExecOp::Clamp, 3, "clamp", 12.0),
(ExecOp::Exp, 1, "exp", 13.0),
(ExecOp::Log, 1, "log", 14.0),
(ExecOp::Sin, 1, "sin", 15.0),
(ExecOp::Cos, 1, "cos", 16.0),
(ExecOp::Tanh, 1, "tanh", 17.0),
(ExecOp::Sqrt, 1, "sqrt", 18.0),
(ExecOp::Rsqrt, 1, "rsqrt", 19.0),
(ExecOp::Pow, 2, "pow", 20.0),
(ExecOp::Expm1, 1, "expm1", 21.0),
(ExecOp::Log1p, 1, "log1p", 22.0),
(ExecOp::Gather(gather_config()), 2, "gather", 33.0),
(ExecOp::Scatter(scatter_config()), 3, "scatter", 34.0),
(
ExecOp::Slice(SliceConfig {
starts: vec![0],
limits: vec![1],
strides: vec![1],
}),
1,
"slice",
35.0,
),
(
ExecOp::DynamicSlice {
slice_sizes: vec![1],
},
2,
"dynamic_slice",
36.0,
),
(ExecOp::DynamicUpdateSlice, 3, "dynamic_update_slice", 36.5),
(ExecOp::Pad(pad_config()), 1, "pad", 37.0),
(ExecOp::Concatenate { axis: 0 }, 2, "concatenate", 38.0),
(ExecOp::Reverse { axes: vec![0] }, 1, "reverse", 39.0),
(ExecOp::ReduceProd { axes: vec![0] }, 1, "reduce_prod", 29.0),
(ExecOp::ReduceMax { axes: vec![0] }, 1, "reduce_max", 30.0),
(ExecOp::ReduceMin { axes: vec![0] }, 1, "reduce_min", 31.0),
];
for (op, input_count, expected_call, expected_value) in cases {
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let program = single_instruction_program(op, input_count);
let inputs = (0..input_count)
.map(|idx| scalar_tensor(idx as f64 + 1.0))
.collect();
let outputs = executor.eval_exec_ir(&program, inputs).unwrap();
assert_eq!(executor.backend().calls, vec![expected_call]);
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_value(&outputs[0]), expected_value);
}
}
#[test]
fn eval_exec_ir_resolves_dynamic_gather_slice_sizes_from_shape_sources() {
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let program = ExecProgram {
instructions: vec![ExecInstruction {
op: ExecOp::GatherDynamicSliceSizes {
offset_dims: vec![1],
collapsed_slice_dims: vec![0],
start_index_map: vec![0],
index_vector_dim: 1,
slice_sizes: vec![
DimExpr::Const(1),
DimExpr::InputDim {
input_idx: 2,
axis: 1,
},
],
},
input_slots: vec![0, 1, 2],
output_slots: vec![3],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![true, true, true],
}],
input_slots: vec![0, 1, 2],
output_slots: vec![3],
n_slots: 4,
};
let inputs = vec![
f64_tensor(vec![4, 5], vec![0.0; 20]),
f64_tensor(vec![1, 1], vec![0.0]),
f64_tensor(vec![1, 3], vec![0.0; 3]),
];
let outputs = executor.eval_exec_ir(&program, inputs).unwrap();
assert_eq!(executor.backend().calls, vec!["gather"]);
assert_eq!(executor.backend().last_gather_slice_sizes, Some(vec![1, 3]));
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_value(&outputs[0]), 33.0);
}
#[test]
fn eval_exec_ir_materializes_constant_scalars_without_backend_dispatch() {
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let program = ExecProgram {
instructions: vec![ExecInstruction {
op: ExecOp::Constant {
dtype: DType::F64,
bytes: 2.5_f64.to_le_bytes().to_vec(),
},
input_slots: vec![],
output_slots: vec![0],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![],
}],
input_slots: vec![],
output_slots: vec![0],
n_slots: 1,
};
let outputs = executor.eval_exec_ir(&program, vec![]).unwrap();
assert!(executor.backend().calls.is_empty());
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_value(&outputs[0]), 2.5);
}
#[test]
fn eval_exec_ir_materializes_complex_constants() {
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let value = Complex64::new(1.5, -2.0);
let mut bytes = Vec::new();
bytes.extend_from_slice(&value.re.to_le_bytes());
bytes.extend_from_slice(&value.im.to_le_bytes());
let program = ExecProgram {
instructions: vec![ExecInstruction {
op: ExecOp::Constant {
dtype: DType::C64,
bytes,
},
input_slots: vec![],
output_slots: vec![0],
dtype: DType::C64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![],
}],
input_slots: vec![],
output_slots: vec![0],
n_slots: 1,
};
let outputs = executor.eval_exec_ir(&program, vec![]).unwrap();
assert!(executor.backend().calls.is_empty());
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_c64_value(&outputs[0]), value);
}
#[test]
fn eval_exec_ir_propagates_backend_errors() {
let mut executor = GraphExecutor::new(FakeTensorBackend {
calls: Vec::new(),
error_on: Some("add"),
reclaimed: 0,
last_gather_slice_sizes: None,
});
let err = executor
.eval_exec_ir(
&single_instruction_program(ExecOp::Add, 2),
vec![scalar_tensor(1.0), scalar_tensor(2.0)],
)
.unwrap_err();
assert_eq!(executor.backend().calls, vec!["add"]);
assert!(matches!(
err,
Error::TensorRuntime(tenferro_tensor::Error::BackendFailure { op: "add", .. })
));
}
#[test]
fn eval_exec_ir_reports_missing_slots_as_runtime_errors() {
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let program = ExecProgram {
instructions: vec![ExecInstruction {
op: ExecOp::Add,
input_slots: vec![0, 1],
output_slots: vec![2],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![false, false],
}],
input_slots: vec![0],
output_slots: vec![2],
n_slots: 3,
};
let err = executor
.eval_exec_ir(&program, vec![scalar_tensor(1.0)])
.unwrap_err();
assert!(executor.backend().calls.is_empty());
assert!(matches!(
err,
Error::TensorRuntime(tenferro_tensor::Error::MissingValue { slot: 1 })
));
}
#[test]
fn eval_exec_ir_reclaims_last_use_host_buffers() {
let program = ExecProgram {
instructions: vec![
ExecInstruction {
op: ExecOp::Add,
input_slots: vec![0, 1],
output_slots: vec![2],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![true, true],
},
ExecInstruction {
op: ExecOp::Negate,
input_slots: vec![2],
output_slots: vec![3],
dtype: DType::F64,
output_shapes: vec![Vec::new()].into(),
output_extents: empty_extents(1).into(),
last_use: vec![true],
},
],
input_slots: vec![0, 1],
output_slots: vec![3],
n_slots: 4,
};
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let outputs = executor
.eval_exec_ir(&program, vec![scalar_tensor(1.0), scalar_tensor(2.0)])
.unwrap();
assert_eq!(executor.backend().calls, vec!["add", "neg"]);
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_value(&outputs[0]), 3.0);
assert_eq!(executor.backend().reclaimed, 3);
}
#[test]
fn graph_executor_does_not_reclaim_borrowed_input_slots() {
let x = TracedTensor::input_symbolic_shape(DType::F64, 1).unwrap();
let y = (&x + &x).unwrap().neg();
let mut compiler = GraphCompiler::new();
let program = compiler
.compile_with_input_specs(&y, &[(&x, DType::F64, &[2])])
.unwrap();
let input = f64_tensor(vec![2], vec![1.0, 2.0]);
let mut executor = GraphExecutor::new(FakeTensorBackend::default());
let outputs = executor
.run_many_with_input_reads(&program, &[(&x, TensorRead::from_tensor(&input))])
.unwrap();
assert_eq!(executor.backend().calls, vec!["add", "neg"]);
assert_eq!(executor.backend().reclaimed, 1);
assert_eq!(outputs.len(), 1);
assert_eq!(scalar_value(&outputs[0]), 3.0);
assert_eq!(input.as_slice::<f64>().unwrap(), &[1.0, 2.0]);
}