#include <cstdint>
#include <fstream>
#include <iostream>
#include <stdexcept>
#include <vector>

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

int main(int argc, char **argv) try {
  if (argc < 6 || argc > 8)
    return 2;
  const std::size_t items = std::stoul(argv[2]);
  const unsigned sourceTile = std::stoul(argv[3]);
  const unsigned firstDestinationTile = std::stoul(argv[4]);
  const unsigned secondDestinationTile = std::stoul(argv[5]);
  const std::size_t sourcePaddingItems = argc >= 7 ? std::stoul(argv[6]) : 0;
  const std::size_t destinationPaddingItems =
      argc == 8 ? std::stoul(argv[7]) : sourcePaddingItems;
  const std::vector<std::size_t> paddingItems = {
      sourcePaddingItems, destinationPaddingItems, destinationPaddingItems};

  poplar::DeviceManager manager;
  auto devices = manager.getDevices(poplar::TargetType::IPU, 1);
  if (devices.empty() || !devices.front().attach())
    return 3;
  poplar::Graph graph(devices.front().getTarget());
  graph.addCodelets("codelets64.cpp", poplar::CodeletFileType::Auto, "-O2",
                    "ipu2");

  std::vector<poplar::Tensor> padding;
  const std::vector<std::pair<const char *, unsigned>> paddingSpecs = {
      {"source-padding", sourceTile}, {"first-padding", firstDestinationTile},
      {"second-padding", secondDestinationTile}};
  for (std::size_t index = 0; index != paddingSpecs.size(); ++index) {
    if (paddingItems[index] == 0)
      continue;
    const auto [name, tile] = paddingSpecs[index];
    auto tensor = graph.addVariable(poplar::UNSIGNED_LONGLONG,
                                    {paddingItems[index]}, name);
    graph.setTileMapping(tensor, tile);
    graph.createHostWrite(std::string(name) + "-write", tensor);
    padding.push_back(tensor);
  }

  auto source =
      graph.addVariable(poplar::UNSIGNED_LONGLONG, {items}, "source64");
  auto first =
      graph.addVariable(poplar::UNSIGNED_LONGLONG, {items}, "first64");
  auto second =
      graph.addVariable(poplar::UNSIGNED_LONGLONG, {items}, "second64");
  graph.setTileMapping(source, sourceTile);
  graph.setTileMapping(first, firstDestinationTile);
  graph.setTileMapping(second, secondDestinationTile);

  graph.createHostWrite("source-write", source);
  graph.createHostRead("first-read", first);
  graph.createHostRead("second-read", second);
  const std::vector<unsigned> observedTiles = {
      sourceTile, sourceTile ^ 1u, firstDestinationTile,
      secondDestinationTile};
  auto beforeConfig = graph.addVariable(
      poplar::UNSIGNED_INT, {observedTiles.size(), 11}, "before-config");
  auto afterConfig = graph.addVariable(
      poplar::UNSIGNED_INT, {observedTiles.size(), 11}, "after-config");
  auto before = graph.addComputeSet("read-exchange-config-before");
  auto after = graph.addComputeSet("read-exchange-config-after");
  for (std::size_t index = 0; index != observedTiles.size(); ++index) {
    graph.setTileMapping(beforeConfig[index], observedTiles[index]);
    graph.setTileMapping(afterConfig[index], observedTiles[index]);
    auto beforeVertex = graph.addVertex(before, "ReadExchangeConfig");
    graph.connect(beforeVertex["values"], beforeConfig[index]);
    graph.setTileMapping(beforeVertex, observedTiles[index]);
    auto afterVertex = graph.addVertex(after, "ReadExchangeConfig");
    graph.connect(afterVertex["values"], afterConfig[index]);
    graph.setTileMapping(afterVertex, observedTiles[index]);
  }
  graph.createHostRead("before-config-read", beforeConfig);
  graph.createHostRead("after-config-read", afterConfig);
  auto keep = graph.addComputeSet("keep-destinations-separated");
  auto vertex = graph.addVertex(keep, "KeepPair64Separated");
  graph.connect(vertex["first"], first);
  graph.connect(vertex["second"], second);
  graph.setTileMapping(vertex, firstDestinationTile);
  for (std::size_t index = 0; index != padding.size(); ++index) {
    auto paddingVertex = graph.addVertex(keep, "KeepPadding64");
    graph.connect(paddingVertex["values"], padding[index]);
    graph.setTileMapping(paddingVertex, observedTiles[index == 0 ? 0 : index + 1]);
  }

  poplar::program::Sequence program;
  // Present the two equal destinations as one broadcast copy. Two sequential
  // Copy operations cannot expose double-width exchange to Poplar's exchange
  // planner because they belong to different exchange epochs.
  program.add(poplar::program::Execute(before));
  program.add(poplar::program::Copy(poplar::concat({source, source}),
                                    poplar::concat({first, second})));
  program.add(poplar::program::Execute(after));
  program.add(poplar::program::Execute(keep));
  poplar::Engine engine(graph, program);
  std::ofstream executable(argv[1], std::ios::binary);
  engine.serializeExecutable(executable);
  executable.close();
  engine.load(devices.front());
  for (std::size_t index = 0; index != paddingSpecs.size(); ++index) {
    if (paddingItems[index] == 0)
      continue;
    std::vector<std::uint64_t> paddingValues(paddingItems[index], 0);
    engine.writeTensor(std::string(paddingSpecs[index].first) + "-write",
                       paddingValues.data(),
                       paddingValues.data() + paddingValues.size());
  }
  std::vector<std::uint64_t> input(items), firstResult(items), secondResult(items);
  for (std::size_t index = 0; index != items; ++index)
    input[index] = 0x6400000000000000ull ^
                   (index * 0x9e3779b97f4a7c15ull);
  engine.writeTensor("source-write", input.data(), input.data() + input.size());
  engine.run(0);
  engine.readTensor("first-read", firstResult.data(),
                    firstResult.data() + firstResult.size());
  engine.readTensor("second-read", secondResult.data(),
                    secondResult.data() + secondResult.size());
  std::vector<std::uint32_t> beforeValues(observedTiles.size() * 11);
  std::vector<std::uint32_t> afterValues(observedTiles.size() * 11);
  engine.readTensor("before-config-read", beforeValues.data(),
                    beforeValues.data() + beforeValues.size());
  engine.readTensor("after-config-read", afterValues.data(),
                    afterValues.data() + afterValues.size());
  if (firstResult != input || secondResult != input)
    throw std::runtime_error("paired 64-bit exchange result mismatch");
  std::cout << "paired64 hardware=PASS items=" << items << '\n';
  for (std::size_t index = 0; index != observedTiles.size(); ++index) {
    std::cout << "tile=" << observedTiles[index] << " before=";
    for (std::size_t field = 0; field != 11; ++field)
      std::cout << (field == 0 ? "" : ",") << beforeValues[index * 11 + field];
    std::cout << " after=";
    for (std::size_t field = 0; field != 11; ++field)
      std::cout << (field == 0 ? "" : ",") << afterValues[index * 11 + field];
    std::cout << '\n';
  }
  return 0;
} catch (const std::exception &error) {
  std::cerr << error.what() << '\n';
  return 1;
}
