#include <poplar/Graph.hpp>
#include <poplar/Target.hpp>
#include <poplin/MatMul.hpp>
#include <poplin/codelets.hpp>
#include <popnn/NonLinearity.hpp>
#include <popnn/codelets.hpp>
#include <popops/codelets.hpp>
#include <poputil/TileMapping.hpp>
#include <fstream>
#include <iostream>

int main() {
  poplar::Graph graph(poplar::Target::createIPUTarget(1, 1472, "C600"));
  poplin::addCodelets(graph);
  popnn::addCodelets(graph);
  popops::addCodelets(graph);
  poplin::PlanningCache cache;
  poplar::OptionFlags options{{"partialsType", "half"},
                             {"disableSRForAMPVertices", "true"}};
  auto input = graph.addVariable(poplar::HALF, {729, 1152}, "input");
  poputil::mapTensorLinearly(graph, input);
  auto up = poplin::createMatMulInputRHS(graph, poplar::HALF, poplar::HALF,
      {729,1152}, {1152,4304}, "upWeights", options, &cache);
  auto down = poplin::createMatMulInputRHS(graph, poplar::HALF, poplar::HALF,
      {729,4304}, {4304,1152}, "downWeights", options, &cache);
  poplar::program::Sequence body;
  auto activation = poplin::matMul(graph, input, up, body, "up", options, &cache);
  popnn::geluInPlace(graph, activation, body, "gelu");
  auto output = poplin::matMul(graph, activation, down, body, "down", options, &cache);
  std::cout << "UP\n";
  poplin::matMulReportPlan(std::cout, graph, poplar::HALF, poplar::HALF,
      {729,1152}, {1152,4304}, options, &cache);
  std::cout << "DOWN\n";
  poplin::matMulReportPlan(std::cout, graph, poplar::HALF, poplar::HALF,
      {729,4304}, {4304,1152}, options, &cache);
  auto upLhs = poplin::createMatMulInputLHS(graph, poplar::HALF, poplar::HALF,
      {729,1152}, {1152,4304}, "upLhs", options, &cache);
  auto downLhs = poplin::createMatMulInputLHS(graph, poplar::HALF, poplar::HALF,
      {729,4304}, {4304,1152}, "downLhs", options, &cache);
  std::ofstream csv("artifacts/sdk-mlp-20260922/mappings.csv");
  csv << "tensor,tile,begin,end\n";
  for (auto entry : {std::make_pair("input",input), {"upWeights",up},
       {"downWeights",down}, {"hidden",activation}, {"output",output},
       {"upPreferredInput",upLhs}, {"downPreferredInput",downLhs}}) {
    auto mapping = graph.getTileMapping(entry.second.flatten());
    for (unsigned tile=0; tile<mapping.size(); ++tile)
      for (auto interval : mapping[tile])
        csv << entry.first << ',' << tile << ',' << interval.begin() << ',' << interval.end() << '\n';
  }
}
