#include "utils.h"

#include "data.h"

#include <cmath>
#include <iostream>
#include <stdexcept>

void set_seed(int seed) {
  torch::manual_seed(seed);
  if (torch::cuda::is_available()) torch::cuda::manual_seed_all(seed);
}

torch::Device get_device(DevicePreference preference) {
  if (preference == DevicePreference::CPU) {
    std::cout << "Using CPU device" << std::endl;
    return torch::kCPU;
  }
  if (torch::cuda::is_available()) {
    std::cout << "Using CUDA device" << std::endl;
    return torch::kCUDA;
  }
  if (preference == DevicePreference::CUDA) {
    throw std::runtime_error(
        "CUDA GPU is required for training, but LibTorch cannot see one. Check that an "
        "NVIDIA driver is installed and running, that nvidia-smi works, and that this "
        "binary is linked against a CUDA-enabled LibTorch build.");
  }
  std::cout << "CUDA unavailable; falling back to CPU device" << std::endl;
  return torch::kCPU;
}

torch::Tensor spherical_pose(float azimuth, float elevation, float radius) {
  const float phi = elevation * (M_PI / 180.0f);
  const float theta = azimuth * (M_PI / 180.0f);
  const float cp = std::cos(phi), sp = std::sin(phi);
  const float ct = std::cos(theta), st = std::sin(theta);

  auto translate = torch::tensor({{1.0f, 0.0f, 0.0f, 0.0f},
                                  {0.0f, 1.0f, 0.0f, 0.0f},
                                  {0.0f, 0.0f, 1.0f, radius},
                                  {0.0f, 0.0f, 0.0f, 1.0f}});
  auto rotate_phi = torch::tensor({{1.0f, 0.0f, 0.0f, 0.0f},
                                   {0.0f, cp, -sp, 0.0f},
                                   {0.0f, sp, cp, 0.0f},
                                   {0.0f, 0.0f, 0.0f, 1.0f}});
  auto rotate_theta = torch::tensor({{ct, 0.0f, -st, 0.0f},
                                     {0.0f, 1.0f, 0.0f, 0.0f},
                                     {st, 0.0f, ct, 0.0f},
                                     {0.0f, 0.0f, 0.0f, 1.0f}});
  // Axis flip into the Blender convention: camera looks down -z with +y up.
  auto flip = torch::tensor({{-1.0f, 0.0f, 0.0f, 0.0f},
                             {0.0f, 0.0f, 1.0f, 0.0f},
                             {0.0f, 1.0f, 0.0f, 0.0f},
                             {0.0f, 0.0f, 0.0f, 1.0f}});
  return flip.matmul(rotate_theta.matmul(rotate_phi.matmul(translate)));
}

void render_views(const NeRFRenderer &renderer, const std::string &prefix, int H, int W,
                  int n_frames, const std::filesystem::path &out_dir, float radius,
                  const RenderOptions &opt) {
  std::cout << "Saving " << n_frames << " sample views..." << std::endl;
  for (int i = 0; i < n_frames; i++) {
    const auto azimuth = static_cast<float>(i) * 360.0f / static_cast<float>(n_frames);
    auto pose = spherical_pose(azimuth, -30.0f, radius).to(renderer.device());
    auto out = renderer.render_image(H, W, pose, opt);

    const std::string suffix = prefix + "_" + std::to_string(i) + ".png";
    save_image(out.rgb, out_dir / ("frame_" + suffix));
    // Fixed-range normalisation keeps the depth scale stable across frames; near
    // surfaces come out bright.
    auto depth = 1.0 - ((out.depth - opt.z_near) / (opt.z_far - opt.z_near)).clamp(0.0, 1.0);
    save_image(depth.unsqueeze(-1).expand({H, W, 3}), out_dir / ("frame_depth_" + suffix));
  }
}

void save_checkpoint(const std::filesystem::path &path, const torch::nn::Module &model,
                     int iter) {
  torch::serialize::OutputArchive archive;
  model.save(archive);
  archive.write("epoch", iter);
  archive.save_to(path.string());
  std::cout << "Model weights saved" << std::endl;
}

int load_checkpoint(const std::filesystem::path &path, torch::nn::Module &model) {
  if (!std::filesystem::exists(path))
    throw std::runtime_error("No checkpoint at '" + path.string() + "'");

  torch::serialize::InputArchive archive;
  archive.load_from(path.string());
  model.load(archive);

  c10::IValue iter;
  archive.read("epoch", iter);
  std::cout << "Loaded " << path << " from iteration " << iter.toInt() << std::endl;
  return static_cast<int>(iter.toInt());
}
