#include "config.h"
#include "data.h"
#include "eval.h"
#include "model.h"
#include "renderer.h"
#include "utils.h"

#include <algorithm>
#include <exception>
#include <fstream>
#include <iostream>
#include <random>
#include <vector>

namespace {

torch::Tensor photometric_loss(const torch::Tensor &rgb, const torch::Tensor &target,
                               const Config &cfg) {
  if (cfg.loss == LossType::MSE) return torch::mse_loss(rgb, target);
  // Charbonnier / pseudo-Huber: mean sqrt(c^2 + e^2) - c, which damps the gradient
  // on outlier pixels relative to MSE.
  auto e = rgb - target;
  return (torch::sqrt(cfg.huber_c * cfg.huber_c + e * e) - cfg.huber_c).mean();
}

// Every 8th frame is held out and never trained on.
std::pair<std::vector<int64_t>, std::vector<int64_t>> split_views(int64_t n) {
  std::pair<std::vector<int64_t>, std::vector<int64_t>> split;
  for (int64_t i = 0; i < n; i++) (i % 8 == 0 ? split.second : split.first).push_back(i);
  return split;
}

}  // namespace

int main(int argc, char *argv[]) {
  Config cfg;
  bool help = false;
  if (!parse_arguments(argc, argv, cfg, help)) return help ? 0 : 1;

  set_seed(cfg.seed);

  torch::Device device(torch::kCPU);
  Dataset dataset;
  try {
    device = get_device(cfg.device_pref);
    std::filesystem::create_directories(cfg.output_path);
    dataset = load_dataset(cfg.data_path / "transforms.json", cfg.image_size);
  } catch (const std::exception &e) {
    std::cerr << "Error: " << e.what() << std::endl;
    return 1;
  }

  // Matmul shapes are fixed across iterations, so TF32 and cuDNN autotune are a
  // sizeable throughput win here with no meaningful quality impact.
  if (device.is_cuda()) {
    torch::globalContext().setBenchmarkCuDNN(true);
    torch::globalContext().setAllowTF32CuBLAS(true);
    torch::globalContext().setAllowTF32CuDNN(true);
  }

  // Move the dataset once so per-iteration views avoid host-to-device copies.
  dataset.images = dataset.images.to(device);
  dataset.poses = dataset.poses.to(device);
  const auto [train_idx, test_idx] = split_views(dataset.images.size(0));
  const int H = dataset.images.size(1), W = dataset.images.size(2);

  std::cout << "Images: " << dataset.images.sizes() << "\nFocal length: " << dataset.focal
            << "\nSampler: " << sampler_name(cfg.sampler)
            << "  samples/ray=" << cfg.samples() << "+" << cfg.importance() << std::endl;
  if (cfg.render_checkpoint.empty()) {
    std::cout << "Train views: " << train_idx.size()
              << "  Test views: " << test_idx.size()
              << "  rays/step=" << (cfg.ray_batched() ? cfg.ray_batch : H * W)
              << "  warmup=" << cfg.warmup_iters << std::endl;
  }

  // The model owns the proposal network as a submodule, so one optimiser trains both.
  SirenNeRF model(device, cfg.width, cfg.depth);
  torch::optim::AdamW optimizer(
      model.parameters(),
      torch::optim::AdamWOptions(cfg.learning_rate).weight_decay(cfg.weight_decay));
  NeRFRenderer renderer(model, dataset.focal, device);

  const RenderOptions train_opt{.strategy = cfg.sampler,
                                .z_near = cfg.z_near,
                                .z_far = cfg.z_far,
                                .n_samples = cfg.samples(),
                                .n_importance = cfg.importance(),
                                .batch_size = cfg.batch_size};
  RenderOptions eval_opt = train_opt;
  eval_opt.deterministic = cfg.hierarchical();
  RenderOptions preview_opt = train_opt;
  preview_opt.z_near = 0.8f;
  preview_opt.z_far = 3.2f;
  preview_opt.deterministic = true;
  // Previews render far more points per pass than a training step, so cap the chunk
  // size; --batch-size can still lower it further to fit memory.
  preview_opt.batch_size = std::min(cfg.batch_size, 320000);

  if (!cfg.render_checkpoint.empty()) {
    torch::NoGradGuard no_grad;
    try {
      load_checkpoint(cfg.render_checkpoint, model);
    } catch (const std::exception &e) {
      std::cerr << "Error: " << e.what() << std::endl;
      return 1;
    }
    render_views(renderer, "final", 240, 240, cfg.n_final_frames, cfg.output_path, 2.1f,
                 preview_opt);
    return 0;
  }

  // Pool every training ray up front; each step samples cfg.ray_batch of them.
  torch::Tensor rays_o_all, rays_d_all, rgb_all;
  if (cfg.ray_batched()) {
    std::vector<torch::Tensor> origins, directions, colors;
    for (int64_t idx : train_idx) {
      auto [o, d] = renderer.get_rays(H, W, dataset.poses[idx]);
      origins.push_back(o);
      directions.push_back(d);
      colors.push_back(dataset.images[idx].reshape({-1, 3}));
    }
    rays_o_all = torch::cat(origins, 0);
    rays_d_all = torch::cat(directions, 0);
    rgb_all = torch::cat(colors, 0);
  }

  std::mt19937 view_rng(cfg.seed);
  std::uniform_int_distribution<size_t> view_dist(0, train_idx.size() - 1);

  for (int i = 0; i <= cfg.n_iters; i++) {
    optimizer.zero_grad();

    RenderOutput out;
    torch::Tensor target, rays_o, rays_d;
    if (cfg.ray_batched()) {
      auto sel = torch::randint(0, rays_o_all.size(0), {cfg.ray_batch},
                                torch::dtype(torch::kLong).device(device));
      rays_o = rays_o_all.index_select(0, sel);
      rays_d = rays_d_all.index_select(0, sel);
      target = rgb_all.index_select(0, sel);

      RenderOptions opt = train_opt;
      opt.use_proposal = !(cfg.hierarchical() && i < cfg.warmup_iters);
      out = renderer.render(rays_o, rays_d, opt);
    } else {
      const int64_t idx = train_idx[view_dist(view_rng)];
      target = dataset.images[idx].reshape({-1, 3});
      std::tie(rays_o, rays_d) = renderer.get_rays(H, W, dataset.poses[idx]);
      out = renderer.render(rays_o, rays_d, train_opt);
    }

    auto rgb_loss = photometric_loss(out.rgb, target, cfg);
    auto loss = rgb_loss;
    if (cfg.hierarchical()) {
      // Interlevel loss (mip-NeRF 360): the coarse proposal histogram must upper-
      // bound the final main-network histogram at different bin locations.
      auto w_prop = renderer.proposal_weights(rays_o, rays_d, out.coarse_z, cfg.batch_size);
      loss = loss + cfg.interlevel_weight *
                        interlevel_loss(out.fine_z, out.weights, out.coarse_z, w_prop);
    }
    loss.backward();
    optimizer.step();

    if (i % cfg.log_freq == 0)
      std::cout << "Iteration: " << i << " Loss: " << rgb_loss.item<float>() << std::endl;

    if (i % cfg.eval_freq == 0 && i > 0) {
      torch::NoGradGuard no_grad;
      std::vector<EvalMetrics> per_view;
      for (int64_t idx : test_idx) {
        auto rendered = renderer.render_image(H, W, dataset.poses[idx], eval_opt);
        per_view.push_back(compute_metrics(rendered.rgb, dataset.images[idx]));
      }
      const auto avg = mean_metrics(per_view);
      std::cout << "[Eval iter " << i << "] PSNR=" << avg.psnr << "  RMSE=" << avg.rmse
                << "  SSIM=" << avg.ssim << std::endl;
      write_metrics_csv(cfg.output_path / "eval_metrics.csv", i, per_view);

      if (i == cfg.n_iters) {
        std::ofstream summary(cfg.output_path / "final_metrics.txt");
        summary << "iter " << i << "\nPSNR " << avg.psnr << "\nRMSE " << avg.rmse
                << "\nSSIM " << avg.ssim << "\n";
      }
    }

    if (i % cfg.plot_freq == 0) {
      torch::NoGradGuard no_grad;
      render_views(renderer, std::to_string(i), 240, 240, cfg.n_preview_frames,
                   cfg.output_path, 2.1f, preview_opt);
      save_checkpoint(cfg.output_path / "checkpoint.pt", model, i);
    }
  }

  torch::NoGradGuard no_grad;
  save_checkpoint(cfg.output_path / "checkpoint.pt", model, cfg.n_iters);
  std::cout << "Done" << std::endl;
  render_views(renderer, "final", 240, 240, cfg.n_final_frames, cfg.output_path, 2.1f,
               preview_opt);
  return 0;
}
