// SPDX-License-Identifier: Apache-2.0
// Independent GERO runtime audit. Invokes the actual nntrainer C++ layer.
#include <cmath>
#include <cstring>
#include <iomanip>
#include <iostream>
#include <limits>
#include <vector>
#include <layer_context.h>
#include <pow_layer.h>
#include <var_grad.h>

struct Result {
  std::vector<float> y, dx, repeated_dx;
  bool state_preserved;
};
Result evaluate(const nntrainer::TensorDim &dim, const std::vector<float> &x,
                float exponent, float incoming) {
  nntrainer::PowLayer layer;
  layer.setProperty({"exponent=" + std::to_string(exponent)});
  nntrainer::InitLayerContext init({dim}, {true}, false, "pow_audit");
  layer.finalize(init);
  nntrainer::Var_Grad input(dim, nntrainer::Initializer::NONE, true, true, "input");
  nntrainer::Var_Grad output(dim, nntrainer::Initializer::NONE, true, true, "output");
  nntrainer::RunLayerContext context("pow_audit", true, 0.0f, false,
                                    1.0f, nullptr, false, {}, {&input},
                                    {&output}, {});
  for (size_t i=0;i<x.size();++i) {
    input.getVariableRef().getData<float>()[i]=x[i];
    input.getGradientRef().getData<float>()[i]=42.0f;
    output.getGradientRef().getData<float>()[i]=incoming;
  }
  layer.forwarding(context,true);
  layer.calcDerivative(context);
  Result r; r.state_preserved=true;
  for(size_t i=0;i<x.size();++i) {
    r.y.push_back(output.getVariableRef().getData<float>()[i]);
    r.dx.push_back(input.getGradientRef().getData<float>()[i]);
    input.getGradientRef().getData<float>()[i]=-77.0f;
  }
  layer.calcDerivative(context);
  for(size_t i=0;i<x.size();++i) {
    r.repeated_dx.push_back(input.getGradientRef().getData<float>()[i]);
    r.state_preserved &= std::memcmp(&x[i],input.getVariableRef().getData<float>()+i,sizeof(float))==0;
    r.state_preserved &= std::memcmp(&incoming,output.getGradientRef().getData<float>()+i,sizeof(float))==0;
  }
  return r;
}
void number(float x) {
  if(std::isnan(x))std::cout<<"\"NaN\"";
  else if(std::isinf(x))std::cout<<(x>0?"\"Infinity\"":"\"-Infinity\"");
  else std::cout<<std::setprecision(9)<<x;
}
void array(const std::vector<float> &x) {
  std::cout<<"[";for(size_t i=0;i<x.size();++i){if(i)std::cout<<",";number(x[i]);}std::cout<<"]";
}
int main() {
  const float tiny=std::numeric_limits<float>::denorm_min();
  const float low=std::numeric_limits<float>::min();
  const float high=std::numeric_limits<float>::max();
  const std::vector<float> xs={0.0f,-0.0f,tiny,-tiny,low/4,-low/4,low,-low,-2,2,high,-high};
  const std::vector<nntrainer::TensorDim> shapes={
    nntrainer::TensorDim({1,1,1,12}),nntrainer::TensorDim({2,1,2,3}),nntrainer::TensorDim({3,2,1,2})};
  std::cout<<"{\"zero_exponent\":[";
  bool first=true;
  for(float p:{0.0f,-0.0f})for(size_t s=0;s<shapes.size();++s)for(float g:{-3.0f,0.0f,0.25f,1.0f,3.0f}) {
    auto r=evaluate(shapes[s],xs,p,g);
    if(!first)std::cout<<",";first=false;
    std::cout<<"{\"negative_exponent_zero\":"<<(std::signbit(p)?"true":"false")<<",\"shape_index\":"<<s<<",\"incoming\":";
    number(g);std::cout<<",\"x\":";array(xs);std::cout<<",\"y\":";array(r.y);
    std::cout<<",\"dx\":";array(r.dx);std::cout<<",\"repeated_dx\":";array(r.repeated_dx);
    std::cout<<",\"state_preserved\":"<<(r.state_preserved?"true":"false")<<"}";
  }
  std::cout<<"],\"controls\":[";first=true;
  struct Control{float x,p,y,derivative;};
  const std::vector<Control> controls={{0,1,0,1},{0,2,0,0},{-2,3,-8,12},{4,0.5f,2,0.25f},{2,-1,0.5f,-0.25f}};
  for(const auto &c:controls)for(float g:{-3.0f,0.0f,0.25f,1.0f,3.0f}) {
    auto r=evaluate(nntrainer::TensorDim({1,1,1,1}),{c.x},c.p,g);
    if(!first)std::cout<<",";first=false;
    std::cout<<"{\"x\":";number(c.x);std::cout<<",\"p\":";number(c.p);std::cout<<",\"incoming\":";number(g);
    std::cout<<",\"expected_y\":";number(c.y);std::cout<<",\"expected_dx\":";number(g*c.derivative);
    std::cout<<",\"y\":";number(r.y[0]);std::cout<<",\"dx\":";number(r.dx[0]);
    std::cout<<",\"repeated_dx\":";number(r.repeated_dx[0]);
    std::cout<<",\"state_preserved\":"<<(r.state_preserved?"true":"false")<<"}";
  }
  std::cout<<"],\"finite_difference\":{";
  const float h=1.0f/32.0f;
  auto center=evaluate(nntrainer::TensorDim({1,1,1,1}),{0},0,3);
  auto plus=evaluate(nntrainer::TensorDim({1,1,1,1}),{h},0,3);
  auto minus=evaluate(nntrainer::TensorDim({1,1,1,1}),{-h},0,3);
  std::cout<<"\"h\":";number(h);std::cout<<",\"forward_at_zero\":";number(center.y[0]);
  std::cout<<",\"native_gradient\":";number(center.dx[0]);
  std::cout<<",\"forward_difference\":";number(3*(plus.y[0]-minus.y[0])/(2*h));
  std::cout<<"}}\n";
}
