// Dump JDL plan rows for every sender tile except one receiver tile.

#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 receiverTile = argOr(argv, argc, 1, 1286);
  const unsigned lookupSize = argOr(argv, argc, 2, 3);
  const unsigned elemsPerTile = argOr(argv, argc, 3, 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());
  const unsigned numTiles = graph.getTarget().getNumTiles();
  const unsigned numSenders = numTiles - 1;
  Tensor data = graph.addVariable(INT, {numSenders, elemsPerTile}, "data");
  for (unsigned tile = 0, row = 0; tile < numTiles; ++tile) {
    if (tile == receiverTile) {
      continue;
    }
    graph.setTileMapping(data[row++], tile);
  }

  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 << "receiver=" << receiverTile << " count=" << lookupSize
            << " elemsPerTile=" << elemsPerTile
            << " activeTiles=" << jdlPrograms.numActiveTiles << "\n";
  for (unsigned tile = 0, row = 0; tile < numTiles; ++tile) {
    if (tile == receiverTile) {
      continue;
    }
    std::cout << "sender " << tile << ":";
    for (unsigned col = 0; col < 9; ++col) {
      std::cout << " 0x" << std::hex << words[row * 9 + col] << std::dec;
    }
    std::cout << "\n";
    ++row;
  }
  std::cout << "receiver " << receiverTile << ":";
  const unsigned receiverRow = jdlPrograms.numActiveTiles - 1;
  for (unsigned col = 0; col < 9; ++col) {
    std::cout << " 0x" << std::hex << words[receiverRow * 9 + col] << std::dec;
  }
  std::cout << "\n";
}
