#include <poplar/Graph.hpp>
#include <poplar/Target.hpp>
#include <poplin/MatMul.hpp>
#include <iostream>

int main(int argc, char **argv) {
  if (argc != 7) return 2;
  const auto g = std::stoul(argv[1]), m = std::stoul(argv[2]);
  const auto k = std::stoul(argv[3]), n = std::stoul(argv[4]);
  poplar::Graph graph(poplar::Target::createIPUTarget(1, 1472, "C600"));
  poplar::OptionFlags options{{"partialsType", "half"},
                             {"availableMemoryProportion", argv[6]}};
  auto type = std::string(argv[5]) == "fp8" ? poplar::QUARTER : poplar::HALF;
  poplin::PlanningCache cache;
  const std::vector<std::size_t> a{g, m, k}, b{g, k, n};
  const auto lhs = poplin::createMatMulGroupedInputLHS(
      graph, type, poplar::HALF, a, b, "lhs", options, &cache);
  const auto rhs = poplin::createMatMulGroupedInputRHS(
      graph, type, poplar::HALF, a, b, "rhs", options, &cache);
  std::cout << "operand,tile,begin,end\n";
  for (const auto &entry : {std::make_pair("lhs", lhs), std::make_pair("rhs", rhs)}) {
    const auto mapping = graph.getTileMapping(entry.second.flatten());
    for (unsigned tile = 0; tile < mapping.size(); ++tile)
      for (const auto &interval : mapping[tile])
        std::cout << entry.first << ',' << tile << ',' << interval.begin()
                  << ',' << interval.end() << '\n';
  }
}
