#include "mps.h"
#include "../tensor.h"
#include "../utils.h"
#include <torch/torch.h>

namespace TensorOps {

Napi::Value Mps(Tensor* self, const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();

  if (!torch::mps::is_available()) {
    Napi::Error::New(env, "MPS is not available").ThrowAsJavaScriptException();
    return env.Null();
  }

  torch::Tensor result = self->tensor.to(torch::kMPS);
  return Tensor::NewInstance(env, result);
}

}  // namespace TensorOps
