#include "zero_grad.h"
#include "../tensor.h"
#include <torch/torch.h>

namespace TensorOps {

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

  try {
    // Clear the gradient by setting it to an undefined tensor
    // This is equivalent to PyTorch's tensor.grad = None
    if (self->tensor.grad().defined()) {
      self->tensor.mutable_grad().reset();
    }
    // If grad is already undefined, this is a no-op
  } catch (const std::exception& e) {
    Napi::Error::New(env, e.what()).ThrowAsJavaScriptException();
    return env.Undefined();
  }

  return env.Undefined();
}

} // namespace TensorOps
