use crate::function::UnaryFn;
#[cxx::bridge]
pub mod ffi {
unsafe extern "C++" {
include!("mlx/array.h");
include!("mlx-cxx/transforms.hpp");
#[namespace = "mlx::core"]
type array = crate::array::ffi::array;
#[namespace = "mlx_cxx"]
type CxxUnaryFn;
#[namespace = "mlx_cxx"]
type CxxMultiaryFn;
#[namespace = "mlx_cxx"]
type CxxMultiInputSingleOutputFn;
#[namespace = "mlx_cxx"]
type CxxPairInputSingleOutputFn;
#[namespace = "mlx_cxx"]
type CxxSingleInputPairOutputFn;
#[namespace = "mlx::core"]
#[cxx_name = "ValueAndGradFn"]
type CxxValueAndGradFn;
#[namespace = "mlx::core"]
#[cxx_name = "SimpleValueAndGradFn"]
type CxxSimpleValueAndGradFn;
#[namespace = "mlx_cxx"]
fn simplify(outputs: &[UniquePtr<array>]);
#[namespace = "mlx_cxx"]
fn eval(outputs: &[UniquePtr<array>]);
#[namespace = "mlx_cxx"]
#[rust_name = "vjp_multiary_cxx_fn"]
fn vjp(f: &CxxMultiaryFn, primals: &[UniquePtr<array>], cotangents: &[UniquePtr<array>]) -> [UniquePtr<CxxVector<array>>; 2];
#[namespace = "mlx_cxx"]
#[rust_name = "vjp_unary_cxx_fn"]
fn vjp(f: &CxxUnaryFn, primal: &array, cotangent: &array) -> [UniquePtr<array>; 2];
#[namespace = "mlx_cxx"]
#[rust_name = "jvp_multiary_cxx_fn"]
fn jvp(f: &CxxMultiaryFn, primals: &[UniquePtr<array>], tangents: &[UniquePtr<array>]) -> [UniquePtr<CxxVector<array>>; 2];
#[namespace = "mlx_cxx"]
#[rust_name = "jvp_unary_cxx_fn"]
fn jvp(f: &CxxUnaryFn, primal: &array, tangent: &array) -> [UniquePtr<array>; 2];
#[namespace = "mlx_cxx"]
#[rust_name = "value_and_grad_multiary_cxx_fn_argnums"]
fn value_and_grad(f: &CxxMultiaryFn, argnums: &CxxVector<i32>) -> UniquePtr<CxxValueAndGradFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "value_and_grad_multiary_cxx_fn_argnum"]
fn value_and_grad(f: &CxxMultiaryFn, argnum: i32) -> UniquePtr<CxxValueAndGradFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "value_and_grad_unary_cxx_fn"]
fn value_and_grad(f: &CxxUnaryFn) -> UniquePtr<CxxSingleInputPairOutputFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "value_and_grad_multi_input_single_output_cxx_fn"]
fn value_and_grad(f: &CxxMultiInputSingleOutputFn, argnums: &CxxVector<i32>) -> UniquePtr<CxxSimpleValueAndGradFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "grad_multi_input_single_output_cxx_fn_argnums"]
fn grad(f: &CxxMultiInputSingleOutputFn, argnums: &CxxVector<i32>) -> UniquePtr<CxxMultiaryFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "grad_multi_input_single_output_cxx_fn_argnum"]
fn grad(f: &CxxMultiInputSingleOutputFn, argnum: i32) -> UniquePtr<CxxMultiaryFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "grad_unary_cxx_fn"]
fn grad(f: &CxxUnaryFn) -> UniquePtr<CxxUnaryFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "vmap_unary_cxx_fn"]
fn vmap(f: &CxxUnaryFn, in_axis: i32, out_axis: i32) -> UniquePtr<CxxUnaryFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "vmap_pair_input_single_output_cxx_fn"]
fn vmap(f: &CxxPairInputSingleOutputFn, in_axis_a: i32, in_axis_b: i32, out_axis: i32) -> UniquePtr<CxxPairInputSingleOutputFn>;
#[namespace = "mlx_cxx"]
#[rust_name = "vmap_multiary_cxx_fn"]
fn vmap(f: &CxxMultiaryFn, in_axes: &CxxVector<i32>, out_axes: &CxxVector<i32>) -> UniquePtr<CxxMultiaryFn>;
}
}