#include "data.h"

#include "nlohmann/json.hpp"

#define STB_IMAGE_IMPLEMENTATION
#include "stb_image.h"

#define STB_IMAGE_WRITE_IMPLEMENTATION
#include "stb_image_write.h"

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

torch::Tensor load_image(const std::filesystem::path &path) {
  int width = 0, height = 0;
  uint8_t *data = stbi_load(path.string().c_str(), &width, &height, nullptr, 3);
  if (!data) throw std::runtime_error("Could not read image '" + path.string() + "'");

  auto image = torch::from_blob(data, {height, width, 3}, torch::kUInt8).clone();
  stbi_image_free(data);
  return image.to(torch::kFloat32) / 255.0f;
}

void save_image(const torch::Tensor &image, const std::filesystem::path &path) {
  auto bytes = image.mul(255).clamp(0, 255).to(torch::kU8).to(torch::kCPU).contiguous();
  const int height = bytes.size(0), width = bytes.size(1);
  // Warn rather than throw: a failed preview must not abort a long training run.
  if (stbi_write_png(path.string().c_str(), width, height, 3, bytes.data_ptr(),
                     width * 3) == 0) {
    std::cerr << "Failed to save: " << path << std::endl;
  }
}

Dataset load_dataset(const std::filesystem::path &json_path, int target_width) {
  std::ifstream file(json_path);
  if (!file.is_open()) {
    throw std::runtime_error("Could not open '" + json_path.string() +
                             "'. The data path must be a directory containing "
                             "transforms.json and the image files it references.");
  }
  nlohmann::json data;
  file >> data;

  const auto dir = json_path.parent_path();
  std::vector<torch::Tensor> images, poses;
  for (const auto &frame : data["frames"]) {
    images.push_back(load_image(dir / (frame["file_path"].get<std::string>() + ".png")));

    std::array<float, 16> m{};
    for (int i = 0; i < 4; i++)
      for (int j = 0; j < 4; j++) m[i * 4 + j] = frame["transform_matrix"][i][j];
    poses.push_back(torch::from_blob(m.data(), {4, 4}, torch::kFloat32).clone());
  }
  if (images.empty()) throw std::runtime_error("No frames in '" + json_path.string() + "'");

  auto stacked = torch::stack(images);
  const float aspect = static_cast<float>(stacked.size(2)) / stacked.size(1);
  const int target_height = static_cast<int>(target_width * aspect);
  if (stacked.size(1) != target_height || stacked.size(2) != target_width) {
    stacked = torch::nn::functional::interpolate(
                  stacked.permute({0, 3, 1, 2}),  // NHWC -> NCHW
                  torch::nn::functional::InterpolateFuncOptions()
                      .size(std::vector<int64_t>{target_height, target_width})
                      .mode(torch::kBilinear)
                      .align_corners(false))
                  .permute({0, 2, 3, 1});
  }

  const float camera_angle_x = data["camera_angle_x"];
  return {stacked, torch::stack(poses),
          0.5f * static_cast<float>(target_width) / std::tan(0.5f * camera_angle_x)};
}
