#include <cassert>
#include <iostream>
#include "mlx/mlx.h"
using namespace mlx::core;
void array_basics() {
array x(1.0);
auto s = x.item<float>();
assert(s == 1.0);
size_t size = x.size();
assert(size == 1);
int ndim = x.ndim();
assert(ndim == 0);
auto shape = x.shape();
assert(shape.empty());
auto dtype = x.dtype();
assert(dtype == float32);
x = array(1, int32);
assert(x.dtype() == int32);
x.item<int>();
x = array({1.0f, 2.0f, 3.0f, 4.0f}, {2, 2});
auto y = ones({2, 2});
auto z = add(x, y);
z = x + y;
assert(z.dtype() == float32);
assert(z.shape(0) == 2);
assert(z.shape(1) == 2);
eval(z);
eval(z);
z = ones({1});
z.item<float>();
z = ones({2, 2});
std::cout << z << std::endl; }
void automatic_differentiation() {
auto fn = [](array x) { return square(x); };
auto grad_fn = grad(fn);
auto x = array(1.5);
auto dfdx = grad_fn(x);
auto df2dx2 = grad(grad(fn))(x);
}
int main() {
array_basics();
automatic_differentiation();
}