use ferrotorch_core::grad_fns::arithmetic::mul;
use ferrotorch_core::grad_fns::reduction::sum;
use ferrotorch_core::storage::TensorStorage;
use ferrotorch_core::{FerrotorchResult, Tensor};
use ferrotorch_jit::TracedModule;
use ferrotorch_jit_script::script;
fn t1d(data: &[f32]) -> Tensor<f32> {
Tensor::from_storage(TensorStorage::cpu(data.to_vec()), vec![data.len()], false)
.expect("from_storage on the cpu path is infallible for a flat f32 slice")
}
#[script]
fn weighted_sum(a: Tensor<f32>, w: Tensor<f32>) -> FerrotorchResult<Tensor<f32>> {
let prod = mul(&a, &w)?;
sum(&prod)
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let a = t1d(&[1.0, 2.0, 3.0]);
let w = t1d(&[4.0, 5.0, 6.0]);
let module: TracedModule<f32> = weighted_sum(a, w)?;
println!(
"[script_macro_demo] captured TracedModule<f32> (graph nodes invisible via public API)"
);
let a2 = t1d(&[1.0, 2.0, 3.0]);
let w2 = t1d(&[4.0, 5.0, 6.0]);
let result = module.forward_multi(&[a2, w2])?;
let result_data = result.data_vec()?;
let expected = 32.0_f32;
if (result_data[0] - expected).abs() > 1e-5 {
return Err(format!(
"[script_macro_demo] expected {expected}, got {}",
result_data[0]
)
.into());
}
println!(
"[script_macro_demo] forward_multi(weighted_sum) = {} (expected {expected})",
result_data[0]
);
let bytes = module.to_bytes();
let loaded: TracedModule<f32> = TracedModule::<f32>::from_bytes(&bytes)?;
let r = loaded.forward_multi(&[t1d(&[2.0, 3.0]), t1d(&[4.0, 5.0])])?;
let r_data = r.data_vec()?;
let expected_rt = 23.0_f32;
if (r_data[0] - expected_rt).abs() > 1e-5 {
return Err(format!(
"[script_macro_demo] roundtrip expected {expected_rt}, got {}",
r_data[0]
)
.into());
}
println!(
"[script_macro_demo] to_bytes/from_bytes/forward_multi = {} (expected {expected_rt})",
r_data[0]
);
Ok(())
}