// Dump the JIT Dynamic Lookup plan buffer after running the setup vertex.

#include <cstdint>
#include <cstdlib>
#include <iostream>
#include <vector>

#include <poplar/DeviceManager.hpp>
#include <poplar/Engine.hpp>
#include <poplar/Graph.hpp>
#include <poplar/Program.hpp>

#include "JDL.hpp"

using namespace poplar;

namespace {
unsigned argOr(char **argv, int argc, int idx, unsigned fallback) {
  return idx < argc ? static_cast<unsigned>(std::strtoul(argv[idx], nullptr, 0))
                    : fallback;
}
} // namespace

int main(int argc, char **argv) {
  const unsigned senderTile = argOr(argv, argc, 1, 0);
  const unsigned receiverTile = argOr(argv, argc, 2, 1286);
  const unsigned lookupSize = argOr(argv, argc, 3, 3);
  const unsigned elemsPerTile = argOr(argv, argc, 4, 16);

  auto devManager = DeviceManager();
  auto devs = devManager.getDevices(TargetType::IPU, 1);
  if (devs.empty()) {
    std::cerr << "no IPU device available\n";
    return 2;
  }
  Device &device = devs[0];
  if (!device.attach()) {
    std::cerr << "failed to attach device\n";
    return 3;
  }

  Graph graph(device.getTarget());
  Tensor data = graph.addVariable(INT, {1, elemsPerTile}, "data");
  graph.setTileMapping(data[0], senderTile);

  Tensor tileSelector = graph.addVariable(INT, {}, "tileSelector");
  Tensor elementSelector = graph.addVariable(INT, {}, "elementSelector");
  Tensor result = graph.addVariable(INT, {lookupSize}, "result");
  graph.setTileMapping(tileSelector, receiverTile);
  graph.setTileMapping(elementSelector, receiverTile);
  graph.setTileMapping(result, receiverTile);

  JDL::Programs jdlPrograms =
      JDL::createPrograms(graph, data, tileSelector, elementSelector, result);
  graph.createHostRead("planRead", jdlPrograms.planBuf.flatten());

  Engine engine(graph, program::Sequence({jdlPrograms.setup}));
  engine.load(device);
  engine.run(0);

  std::vector<std::uint32_t> words(jdlPrograms.numActiveTiles * 9);
  engine.readTensor("planRead", words.data(), words.data() + words.size());

  std::cout << "sender=" << senderTile << " receiver=" << receiverTile
            << " count=" << lookupSize << " elemsPerTile=" << elemsPerTile
            << " activeTiles=" << jdlPrograms.numActiveTiles << "\n";
  for (unsigned row = 0; row < jdlPrograms.numActiveTiles; ++row) {
    std::cout << "row" << row << ":";
    for (unsigned col = 0; col < 9; ++col) {
      std::cout << " 0x" << std::hex << words[row * 9 + col] << std::dec;
    }
    std::cout << "\n";
  }
}
