#include "exchange_plan.hpp"

#include <algorithm>
#include <cstdint>
#include <cstdlib>
#include <fstream>
#include <iostream>
#include <stdexcept>
#include <string>
#include <vector>

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

using namespace poplar;

namespace {
struct Pair {
  unsigned source;
  unsigned destination;
};
} // namespace

int main() {
  DeviceManager manager;
  auto devices = manager.getDevices(TargetType::IPU, 1);
  if (devices.empty() || !devices[0].attach()) {
    std::cerr << "failed to attach an IPU\n";
    return 2;
  }
  Device &device = devices[0];
  Graph graph(device.getTarget());
  graph.addCodelets(std::vector<std::string>{"no_blob_codelets.cpp"}, "-O2",
                    "ipu2");
  const unsigned tiles = graph.getTarget().getNumTiles();
  unsigned rounds = 0;
  while ((1u << rounds) < tiles) ++rounds;

  Tensor partials = graph.addVariable(INT, {tiles}, "partials");
  Tensor received = graph.addVariable(INT, {rounds, tiles}, "received");
  Tensor plans = graph.addVariable(UNSIGNED_INT, {rounds, tiles, 9}, "plans");
  Tensor dummy = graph.addVariable(UNSIGNED_INT, {rounds, tiles, 1}, "dummy");
  Tensor selectors = graph.addConstant(INT, {rounds, tiles}, 0, "selectors");
  std::vector<int> initial(tiles);
  for (unsigned tile = 0; tile < tiles; ++tile) {
    initial[tile] = static_cast<int>(tile + 1);
    graph.setTileMapping(partials[tile], tile);
    for (unsigned round = 0; round < rounds; ++round) {
      graph.setTileMapping(received[round][tile], tile);
      graph.setTileMapping(plans[round][tile], tile);
      graph.setTileMapping(dummy[round][tile], tile);
      graph.setTileMapping(selectors[round][tile], tile);
    }
  }
  graph.setInitialValue<int>(partials, initial);
  std::vector<int> receivedInitial(received.numElements(), 0);
  graph.setInitialValue<int>(received.flatten(), receivedInitial);

  std::vector<std::uint32_t> planWords(rounds * tiles * 9);
  program::Sequence program;
  for (unsigned round = 0; round < rounds; ++round) {
    const unsigned stride = 1u << round;
    ComputeSet add = graph.addComputeSet("addRound" + std::to_string(round));
    std::vector<Pair> pairs;
    for (unsigned destination = 0; destination < tiles;
         destination += stride * 2) {
      if (destination + stride < tiles)
        pairs.push_back({destination + stride, destination});
    }
    const char *launchWidthText = std::getenv("XCOM_LAUNCH_WIDTH");
    const unsigned launchWidth = launchWidthText
                                     ? static_cast<unsigned>(std::strtoul(
                                           launchWidthText, nullptr, 0))
                                     : static_cast<unsigned>(pairs.size());
    if (launchWidth == 0) {
      throw std::invalid_argument("XCOM_LAUNCH_WIDTH must be positive");
    }
    const unsigned launches = (pairs.size() + launchWidth - 1) / launchWidth;
    std::cout << "round " << round << ": " << pairs.size() << " transfers, "
              << launches << " launches\n";
    for (unsigned launch = 0; launch < launches; ++launch) {
      ComputeSet exchange = graph.addComputeSet(
          "exchangeRound" + std::to_string(round) + "Launch" +
          std::to_string(launch));
      std::vector<bool> active(tiles, false);
      const unsigned begin = launch * launchWidth;
      const unsigned end = std::min<unsigned>(pairs.size(), begin + launchWidth);
      for (unsigned index = begin; index < end; ++index) {
        const Pair &transfer = pairs[index];
        const unsigned source = transfer.source;
        const unsigned destination = transfer.destination;
        active[source] = active[destination] = true;
        std::uint32_t *sourceRow = &planWords[(round * tiles + source) * 9];
        std::uint32_t *destinationRow =
            &planWords[(round * tiles + destination) * 9];
        xcom::Plan pair = xcom::assemble(source, destination, 1);
        std::copy(pair.sender.begin(), pair.sender.end(), sourceRow);
        auto vertex = graph.addVertex(
            exchange, "NoBlobSend",
            {{"planBuf", plans[round][source]},
             {"nonexecutableDummy", dummy[round][source]},
             {"elementSelector", selectors[round][source]},
             {"data", partials.slice(source, source + 1)}});
        graph.setTileMapping(vertex, source);
        pair.receiver[2] ^= xcom::logicalToPhysical(source);
        std::copy(pair.receiver.begin() + 1, pair.receiver.end(), destinationRow);
        auto recv = graph.addVertex(
            exchange, "NoBlobRecvFixed",
            {{"planBuf", plans[round][destination]},
             {"nonexecutableDummy", dummy[round][destination]},
             {"result", received[round].slice(destination, destination + 1)}});
        graph.setTileMapping(recv, destination);
      }
      for (unsigned tile = 0; tile < tiles; ++tile) {
        if (active[tile]) continue;
        auto inactive = graph.addVertex(exchange, "NoBlobNonParticipation");
        graph.setTileMapping(inactive, tile);
      }
      if (std::getenv("XCOM_SYNC_EACH_LAUNCH"))
        program.add(program::Sync(SyncType::INTERNAL));
      program.add(program::Execute(exchange));
    }
    for (const Pair &transfer : pairs) {
      auto sum = graph.addVertex(add, "AddExchangeValue",
                                 {{"partial", partials[transfer.destination]},
                                  {"received", received[round][transfer.destination]}});
      graph.setTileMapping(sum, transfer.destination);
    }
    program.add(program::Execute(add));
  }
  graph.createHostWrite("plans", plans.flatten());
  graph.createHostRead("partials", partials);
  Engine engine(graph, program);
  if (const char *path = std::getenv("XCOM_EXECUTABLE")) {
    std::ofstream output(path, std::ios::binary);
    if (!output) throw std::runtime_error("cannot create executable output");
    engine.serializeExecutable(output);
  }
  engine.load(device);
  engine.writeTensor("plans", planWords.data(), planWords.data() + planWords.size());
  engine.run(0);
  std::vector<int> result(tiles);
  engine.readTensor("partials", result.data(), result.data() + result.size());
  const int expected = static_cast<int>(tiles * (tiles + 1) / 2);
  const bool ok = result[0] == expected;
  std::cout << "tiles=" << tiles << " rounds=" << rounds
            << " sum=" << result[0] << " expected=" << expected << ' '
            << (ok ? "PASS\n" : "FAIL\n");
  return ok ? 0 : 1;
}
