#include "eval.h"

#include <cmath>
#include <fstream>
#include <iostream>
#include <limits>

namespace {

// Separable Gaussian as a depthwise conv kernel: [channels, 1, size, size].
torch::Tensor gaussian_kernel(int size, float sigma, int64_t channels,
                              const torch::TensorOptions &options) {
  auto x = torch::arange(size, options) - (size - 1) / 2.0;
  auto k = torch::exp(-x * x / (2.0f * sigma * sigma));
  k = k / k.sum();
  return torch::outer(k, k).expand({channels, 1, size, size}).contiguous();
}

// Single-scale SSIM with an 11x11 Gaussian window, averaged over channels.
float compute_ssim(const torch::Tensor &rendered, const torch::Tensor &target) {
  constexpr int kWindow = 11;
  constexpr float kC1 = 0.01f * 0.01f, kC2 = 0.03f * 0.03f;  // (K * L)^2 with L = 1
  const int64_t channels = rendered.size(2);
  auto kernel = gaussian_kernel(kWindow, 1.5f, channels, rendered.options());
  auto conv_opts =
      torch::nn::functional::Conv2dFuncOptions().padding(kWindow / 2).groups(channels);
  auto blur = [&](const torch::Tensor &t) {
    return torch::nn::functional::conv2d(t, kernel, conv_opts);
  };

  auto x = rendered.permute({2, 0, 1}).unsqueeze(0).contiguous();
  auto y = target.permute({2, 0, 1}).unsqueeze(0).contiguous();
  auto mx = blur(x), my = blur(y);
  auto mx2 = mx * mx, my2 = my * my, mxy = mx * my;
  auto vx = blur(x * x) - mx2, vy = blur(y * y) - my2, vxy = blur(x * y) - mxy;

  auto ssim = ((2.0f * mxy + kC1) * (2.0f * vxy + kC2)) /
              ((mx2 + my2 + kC1) * (vx + vy + kC2));
  return ssim.mean().item<float>();
}

}  // namespace

EvalMetrics compute_metrics(const torch::Tensor &rendered, const torch::Tensor &target) {
  const float mse = torch::mse_loss(rendered, target).item<float>();
  return {mse > 0.0f ? 10.0f * std::log10(1.0f / mse)
                     : std::numeric_limits<float>::infinity(),
          std::sqrt(mse), compute_ssim(rendered, target)};
}

EvalMetrics mean_metrics(const std::vector<EvalMetrics> &views) {
  EvalMetrics sum{0.0f, 0.0f, 0.0f};
  for (const auto &m : views) {
    sum.psnr += m.psnr;
    sum.rmse += m.rmse;
    sum.ssim += m.ssim;
  }
  const auto n = static_cast<float>(views.size());
  return {sum.psnr / n, sum.rmse / n, sum.ssim / n};
}

void write_metrics_csv(const std::filesystem::path &path, int iter,
                       const std::vector<EvalMetrics> &views) {
  const bool write_header = !std::filesystem::exists(path);
  std::ofstream file(path, std::ios::app);
  // Warn rather than throw: losing a metrics row must not abort a long training run.
  if (!file.is_open()) {
    std::cerr << "Failed to open metrics CSV: " << path << std::endl;
    return;
  }

  if (write_header) file << "iter,view,psnr,rmse,ssim\n";
  for (size_t v = 0; v < views.size(); v++) {
    file << iter << "," << v << "," << views[v].psnr << "," << views[v].rmse << ","
         << views[v].ssim << "\n";
  }
}
