// Standalone SDK FP16 MLP: compact input, resident weights, timed device preparation
// and computation. Host transfers and compilation are outside the cycle counter.
// Build with the SDK enabled:
// c++ -O2 scripts/sdk-mlp-benchmark.cpp -o sdk-mlp-benchmark \
//   -lpoplar -lpoplin -lpopops -lpopnn -lpoputil
#include <poplar/CycleCount.hpp>
#include <poplar/DeviceManager.hpp>
#include <poplar/Engine.hpp>
#include <poplin/MatMul.hpp>
#include <poplin/codelets.hpp>
#include <popnn/NonLinearity.hpp>
#include <popnn/codelets.hpp>
#include <popops/codelets.hpp>
#include <poputil/TileMapping.hpp>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <fstream>
#include <iostream>
#include <random>
#include <stdexcept>
#include <vector>

int main(int argc, char **argv) try {
  if (argc != 2) {
    std::cerr << "usage: sdk-mlp-benchmark OUTPUT.json\n";
    return 2;
  }
  constexpr unsigned rows = 729, width = 1152, hidden = 4304;
  auto manager = poplar::DeviceManager::createDeviceManager();
  auto devices = manager.getDevices(poplar::TargetType::IPU, 1);
  if (devices.empty() || !devices[0].attach())
    throw std::runtime_error("cannot attach IPU");
  const auto &target = devices[0].getTarget();
  poplar::Graph graph(target);
  poplin::addCodelets(graph);
  popops::addCodelets(graph);
  popnn::addCodelets(graph);
  poplin::PlanningCache cache;
  const poplar::OptionFlags options{{"partialsType", "half"},
                                    {"disableSRForAMPVertices", "true"}};
  auto input = graph.addVariable(poplar::HALF, {rows, width}, "input");
  poputil::mapTensorLinearly(graph, input);
  auto up = poplin::createMatMulInputRHS(graph, poplar::HALF, poplar::HALF,
      {rows, width}, {width, hidden}, "upWeights", options, &cache);
  auto down = poplin::createMatMulInputRHS(graph, poplar::HALF, poplar::HALF,
      {rows, hidden}, {hidden, width}, "downWeights", options, &cache);
  graph.createHostWrite("input", input);
  graph.createHostWrite("up", up);
  graph.createHostWrite("down", down);
  poplar::program::Sequence body;
  auto activation = poplin::matMul(graph, input, up, body, "up", options, &cache);
  popnn::geluInPlace(graph, activation, body, "gelu");
  auto output = poplin::matMul(graph, activation, down, body, "down", options, &cache);
  auto cycles = poplar::cycleCount(graph, body, 0, poplar::SyncType::INTERNAL);
  graph.createHostRead("cycles", cycles);
  graph.createHostRead("output", output);
  poplar::Engine engine(graph, body);
  engine.load(devices[0]);
  std::mt19937 random(729);
  std::normal_distribution<float> normal;
  auto upload = [&](const char *name, size_t count, float scale) {
    std::vector<float> values(count);
    for (auto &value : values) value = normal(random) * scale;
    std::vector<uint16_t> encoded(count);
    poplar::copyFloatToDeviceHalf(target, values.data(), encoded.data(), count);
    // Reference the uploaded, rounded values rather than the original floats.
    poplar::copyDeviceHalfToFloat(target, encoded.data(), values.data(), count);
    engine.writeTensor(name, encoded.data(), encoded.data() + count);
    return values;
  };
  auto x = upload("input", rows * width, 1);
  auto a = upload("up", width * hidden, 1 / std::sqrt(float(width)));
  auto b = upload("down", hidden * width, 1 / std::sqrt(float(hidden)));
  engine.run();
  uint32_t stamps[2];
  engine.readTensor("cycles", stamps, stamps + 2);
  const uint64_t elapsed = uint64_t(stamps[0]) | (uint64_t(stamps[1]) << 32);
  std::vector<uint16_t> raw(rows * width);
  std::vector<float> actual(raw.size());
  engine.readTensor("output", raw.data(), raw.data() + raw.size());
  poplar::copyDeviceHalfToFloat(target, raw.data(), actual.data(), actual.size());
  for (float value : actual)
    if (!std::isfinite(value)) throw std::runtime_error("nonfinite output");
  double error = 0;
  for (unsigned row : {0u, 127u, 364u, 728u}) {
    std::vector<double> h(hidden);
    for (unsigned k = 0; k < width; ++k)
      for (unsigned j = 0; j < hidden; ++j)
        h[j] += double(x[row * width + k]) * a[k * hidden + j];
    for (auto &v : h)
      v = 0.5 * v * (1 + std::tanh(0.7978845608028654 * (v + 0.044715 * v * v * v)));
    for (unsigned column = 0; column < width; column += 137) {
      double expected = 0;
      for (unsigned j = 0; j < hidden; ++j) expected += h[j] * b[j * width + column];
      error = std::max(error, std::abs(expected - actual[row * width + column]));
    }
  }
  std::ofstream report(argv[1]);
  report << "{\"precision\":\"fp16\",\"rows\":" << rows
         << ",\"width\":" << width << ",\"hidden\":" << hidden
         << ",\"tiles\":" << target.getNumTiles() << ",\"cycles\":" << elapsed
         << ",\"sampledMaximumAbsoluteError\":" << error << "}\n";
  std::cout << "cycles=" << elapsed << " sampledMaximumAbsoluteError=" << error << '\n';
  if (error > 0.03) throw std::runtime_error("sampled reference mismatch");
} catch (const std::exception &error) {
  std::cerr << error.what() << '\n';
  return 1;
}
