#include "exchange_plan.hpp"

#include <algorithm>
#include <array>
#include <cstdint>
#include <fstream>
#include <iostream>
#include <string>
#include <vector>

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

using namespace poplar;

namespace {
constexpr unsigned dimension = 4;
constexpr unsigned passes = dimension * 2;
constexpr unsigned wordsPerPlan = 9;
constexpr unsigned passSpacing = 256;
constexpr unsigned firstOutputTile = 8;

constexpr std::array<int, dimension * dimension> matrixA = {
    1,  2,  3,  4,
    5,  6,  7,  8,
    -2, 1,  0,  3,
    9,  -1, 2,  1,
};
constexpr std::array<int, dimension * dimension> matrixB = {
    2,  0,  1,  3,
    1,  -1, 2,  0,
    4,  2,  0,  1,
    -3, 5,  1,  2,
};

unsigned outputTile(unsigned row, unsigned column) {
  return firstOutputTile + row * dimension + column;
}
} // namespace

int main(int argc, char **argv) {
  if (argc != 1 && (argc != 3 || std::string(argv[1]) != "--save")) {
    std::cerr << "usage: multi_tile_matmul [--save EXECUTABLE]\n";
    return 2;
  }
  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>{"matmul_codelets.cpp"}, "-O2",
                    "ipu2");
  const unsigned tiles = graph.getTarget().getNumTiles();

  std::vector<std::uint32_t> planWords(tiles * passes * wordsPerPlan, 0);
  std::vector<unsigned> roles(tiles * passes, 0);
  std::vector<unsigned> slots(tiles * passes, 0);
  std::vector<int> sources(tiles * dimension, 0);
  std::vector<int> receivedInitial(tiles * dimension * 2, 0);
  std::vector<unsigned> outputFlags(tiles, 0);

  for (unsigned row = 0; row < dimension; ++row) {
    const unsigned sourceTile = row;
    std::copy_n(matrixA.begin() + row * dimension, dimension,
                sources.begin() + sourceTile * dimension);
    std::vector<unsigned> destinations;
    for (unsigned column = 0; column < dimension; ++column)
      destinations.push_back(outputTile(row, column));
    const xcom::FanOutPlan multicast = xcom::assembleMulticast(
        sourceTile, destinations, dimension, row * passSpacing);
    roles[sourceTile * passes + row] = 1;
    std::copy(multicast.sender.begin(), multicast.sender.end(),
              planWords.begin() +
                  (sourceTile * passes + row) * wordsPerPlan);
    for (unsigned column = 0; column < dimension; ++column) {
      const unsigned tile = destinations[column];
      roles[tile * passes + row] = 2;
      slots[tile * passes + row] = 0;
      std::copy(multicast.receivers[column].begin(),
                multicast.receivers[column].end(),
                planWords.begin() + (tile * passes + row) * wordsPerPlan);
    }
  }

  for (unsigned column = 0; column < dimension; ++column) {
    const unsigned pass = dimension + column;
    const unsigned sourceTile = dimension + column;
    for (unsigned row = 0; row < dimension; ++row)
      sources[sourceTile * dimension + row] =
          matrixB[row * dimension + column];
    std::vector<unsigned> destinations;
    for (unsigned row = 0; row < dimension; ++row)
      destinations.push_back(outputTile(row, column));
    const xcom::FanOutPlan multicast = xcom::assembleMulticast(
        sourceTile, destinations, dimension, pass * passSpacing);
    roles[sourceTile * passes + pass] = 1;
    std::copy(multicast.sender.begin(), multicast.sender.end(),
              planWords.begin() +
                  (sourceTile * passes + pass) * wordsPerPlan);
    for (unsigned row = 0; row < dimension; ++row) {
      const unsigned tile = destinations[row];
      roles[tile * passes + pass] = 2;
      slots[tile * passes + pass] = dimension;
      std::copy(multicast.receivers[row].begin(),
                multicast.receivers[row].end(),
                planWords.begin() + (tile * passes + pass) * wordsPerPlan);
    }
  }
  for (unsigned row = 0; row < dimension; ++row)
    for (unsigned column = 0; column < dimension; ++column)
      outputFlags[outputTile(row, column)] = 1;

  Tensor plansTensor =
      graph.addVariable(UNSIGNED_INT, {tiles, passes, wordsPerPlan}, "plans");
  Tensor rolesTensor =
      graph.addConstant(UNSIGNED_INT, {tiles, passes}, roles.data(), "roles");
  Tensor slotsTensor =
      graph.addConstant(UNSIGNED_INT, {tiles, passes}, slots.data(), "slots");
  Tensor dummy = graph.addVariable(UNSIGNED_INT, {tiles, 1}, "dummy");
  Tensor sourceTensor =
      graph.addVariable(INT, {tiles, dimension}, "sourceRowsAndColumns");
  Tensor received =
      graph.addVariable(INT, {tiles, dimension * 2}, "receivedRowsAndColumns");
  Tensor isOutput = graph.addConstant(UNSIGNED_INT, {tiles}, outputFlags.data(),
                                      "isOutput");
  Tensor results = graph.addVariable(INT, {tiles}, "results");

  graph.setInitialValue<unsigned>(plansTensor, planWords);
  graph.setInitialValue<int>(sourceTensor, sources);
  graph.setInitialValue<int>(received, receivedInitial);
  for (unsigned tile = 0; tile < tiles; ++tile) {
    graph.setTileMapping(plansTensor[tile], tile);
    graph.setTileMapping(rolesTensor[tile], tile);
    graph.setTileMapping(slotsTensor[tile], tile);
    graph.setTileMapping(dummy[tile], tile);
    graph.setTileMapping(sourceTensor[tile], tile);
    graph.setTileMapping(received[tile], tile);
    graph.setTileMapping(isOutput[tile], tile);
    graph.setTileMapping(results[tile], tile);
  }

  ComputeSet matmul = graph.addComputeSet("customExchangeMatMul");
  for (unsigned tile = 0; tile < tiles; ++tile) {
    auto vertex = graph.addVertex(
        matmul, "MultiTileMatMul",
        {{"plans", plansTensor[tile].flatten()},
         {"roles", rolesTensor[tile]},
         {"slots", slotsTensor[tile]},
         {"nonexecutableDummy", dummy[tile]},
         {"source", sourceTensor[tile]},
         {"received", received[tile]},
         {"isOutput", isOutput[tile]},
         {"result", results[tile]}});
    graph.setTileMapping(vertex, tile);
  }
  graph.createHostRead("resultsRead", results);

  Engine engine(graph, program::Execute(matmul));
  if (argc == 3) {
    std::ofstream output(argv[2], std::ios::binary);
    engine.serializeExecutable(output);
    if (!output) {
      std::cerr << "failed to write executable: " << argv[2] << '\n';
      return 3;
    }
  }
  engine.load(device);
  engine.run(0);
  std::vector<int> resultHost(tiles);
  engine.readTensor("resultsRead", resultHost.data(),
                    resultHost.data() + resultHost.size());

  bool ok = true;
  std::cout << "C =\n";
  for (unsigned row = 0; row < dimension; ++row) {
    for (unsigned column = 0; column < dimension; ++column) {
      int expected = 0;
      for (unsigned k = 0; k < dimension; ++k)
        expected += matrixA[row * dimension + k] *
                    matrixB[k * dimension + column];
      const int actual = resultHost[outputTile(row, column)];
      ok &= actual == expected;
      std::cout << actual << (column + 1 == dimension ? '\n' : ' ');
      if (actual != expected)
        std::cerr << "mismatch C[" << row << "][" << column
                  << "]: actual=" << actual << " expected=" << expected
                  << '\n';
    }
  }
  std::cout << "vertexTypes=1 instances=" << tiles
            << " computeSets=1 customPasses=" << passes
            << " sdkMatrixCopies=0 " << (ok ? "PASS\n" : "FAIL\n");
  return ok ? 0 : 1;
}
