#include "renderer.h"

#include <algorithm>
#include <tuple>
#include <vector>

using namespace torch::indexing;

namespace {

torch::Tensor head(const torch::Tensor &t) { return t.index({"...", Slice(None, 1)}); }
torch::Tensor drop_last(const torch::Tensor &t) { return t.index({"...", Slice(None, -1)}); }
torch::Tensor drop_first(const torch::Tensor &t) { return t.index({"...", Slice(1, None)}); }

// Alpha compositing along the last axis: the shared core of both render passes.
torch::Tensor weights_from_sigma(const torch::Tensor &sigma, const torch::Tensor &z_vals) {
  auto dists = torch::cat(
      {drop_first(z_vals) - drop_last(z_vals), torch::full_like(head(z_vals), 1e10)}, -1);
  auto alpha = 1.0 - torch::exp(-sigma * dists);
  auto transmittance = torch::cumprod(1.0 - alpha + 1e-10, -1);
  return alpha * torch::cat({torch::ones_like(head(alpha)), drop_last(transmittance)}, -1);
}

// Midpoints between adjacent samples, used as interior bin edges.
torch::Tensor midpoints(const torch::Tensor &z) {
  return 0.5 * (drop_first(z) + drop_last(z));
}

}  // namespace

NeRFRenderer::NeRFRenderer(SirenNeRF &model, float focal, torch::Device device,
                           torch::Tensor bg_color)
    : model_(model), device_(device), focal_(focal), bg_color_(bg_color.to(device)) {}

std::pair<torch::Tensor, torch::Tensor> NeRFRenderer::get_rays(
    int H, int W, const torch::Tensor &pose) const {
  auto opts = torch::dtype(torch::kFloat32).device(device_);
  auto grid = torch::meshgrid({torch::arange(W, opts), torch::arange(H, opts)}, "xy");

  // Pixel centres on the image plane; y and z are negated because the camera looks
  // down -z with y up (Blender/OpenGL convention).
  auto dirs = torch::stack({(grid[0] - W * 0.5f) / focal_, -(grid[1] - H * 0.5f) / focal_,
                            -torch::ones_like(grid[0])},
                           -1);
  auto rays_d = (dirs.unsqueeze(-2) * pose.index({Slice(0, 3), Slice(0, 3)})).sum(-1);
  auto rays_o = pose.index({Slice(0, 3), -1}).expand(rays_d.sizes());
  return {rays_o.reshape({-1, 3}), rays_d.reshape({-1, 3})};
}

torch::Tensor NeRFRenderer::sample_z(int64_t n_rays, int n_samples,
                                     const RenderOptions &opt, bool jitter) const {
  auto z = torch::linspace(opt.z_near, opt.z_far, n_samples, device_)
               .expand({n_rays, n_samples})
               .contiguous();
  if (!jitter) return z;
  // NeRF 5.2: split the interval into bins and draw one sample uniformly per bin.
  auto mids = midpoints(z);
  auto upper = torch::cat({mids, z.index({"...", Slice(-1, None)})}, -1);
  auto lower = torch::cat({head(z), mids}, -1);
  return lower + (upper - lower) * torch::rand_like(z);
}

RenderOutput NeRFRenderer::render(const torch::Tensor &rays_o, const torch::Tensor &rays_d,
                                  const RenderOptions &opt) const {
  const bool jitter = opt.strategy == SampleStrategy::STRATIFIED ||
                      (opt.strategy == SampleStrategy::PROPOSAL && !opt.deterministic);
  auto z_coarse = sample_z(rays_o.size(0), opt.n_samples, opt, jitter);

  if (opt.strategy != SampleStrategy::PROPOSAL || opt.n_importance <= 0) {
    auto out = volume_render(rays_o, rays_d, z_coarse, opt.batch_size);
    out.fine_z = z_coarse;
    return out;
  }

  // The coarse pass only decides where to place fine samples, so it needs no
  // gradients. use_proposal=false reuses the full model instead (warm-up).
  torch::Tensor z_all;
  {
    torch::NoGradGuard no_grad;
    auto w = opt.use_proposal
                 ? proposal_weights(rays_o, rays_d, z_coarse, opt.batch_size)
                 : volume_render(rays_o, rays_d, z_coarse, opt.batch_size).weights;
    auto z_fine = sample_pdf(midpoints(z_coarse), w.index({"...", Slice(1, -1)}),
                             opt.n_importance, opt.deterministic);
    z_all = std::get<0>(torch::sort(torch::cat({z_coarse, z_fine}, -1), -1));
  }

  auto out = volume_render(rays_o, rays_d, z_all, opt.batch_size);
  out.fine_z = z_all;
  out.coarse_z = z_coarse;
  return out;
}

RenderOutput NeRFRenderer::render_image(int H, int W, const torch::Tensor &pose,
                                        const RenderOptions &opt) const {
  auto [rays_o, rays_d] = get_rays(H, W, pose);
  auto out = render(rays_o, rays_d, opt);
  out.rgb = out.rgb.view({H, W, 3});
  out.depth = out.depth.view({H, W});
  return out;
}

torch::Tensor NeRFRenderer::proposal_weights(const torch::Tensor &rays_o,
                                             const torch::Tensor &rays_d,
                                             const torch::Tensor &z_vals,
                                             int batch_size) const {
  auto pts = (rays_o.unsqueeze(-2) + rays_d.unsqueeze(-2) * z_vals.unsqueeze(-1))
                 .reshape({-1, 3});
  std::vector<torch::Tensor> chunks;
  for (int64_t i = 0; i < pts.size(0); i += batch_size) {
    const auto end = std::min<int64_t>(i + batch_size, pts.size(0));
    chunks.push_back(model_.proposal_sigma(pts.slice(0, i, end)));
  }
  return weights_from_sigma(torch::cat(chunks, 0).view(z_vals.sizes()), z_vals);
}

RenderOutput NeRFRenderer::volume_render(const torch::Tensor &rays_o,
                                         const torch::Tensor &rays_d,
                                         const torch::Tensor &z_vals,
                                         int batch_size) const {
  auto pts = rays_o.unsqueeze(-2) + rays_d.unsqueeze(-2) * z_vals.unsqueeze(-1);
  auto view_dirs = rays_d / (rays_d.norm(2, -1, true) + 1e-8);
  auto pts_flat = pts.reshape({-1, 3});
  auto dirs_flat = view_dirs.unsqueeze(-2).expand(pts.sizes()).reshape({-1, 3});

  std::vector<torch::Tensor> rgb_chunks, sigma_chunks;
  for (int64_t i = 0; i < pts_flat.size(0); i += batch_size) {
    const auto end = std::min<int64_t>(i + batch_size, pts_flat.size(0));
    auto chunk = model_.forward(pts_flat.slice(0, i, end), dirs_flat.slice(0, i, end));
    rgb_chunks.push_back(chunk.rgb);
    sigma_chunks.push_back(chunk.sigma);
  }

  const int64_t N = z_vals.size(0), S = z_vals.size(1);
  auto rgb = torch::cat(rgb_chunks, 0).view({N, S, 3});
  auto sigma = torch::cat(sigma_chunks, 0).view({N, S});
  auto weights = weights_from_sigma(sigma, z_vals);

  RenderOutput out;
  out.rgb = (weights.unsqueeze(-1) * rgb).sum(-2) +
            (1.0 - weights.sum(-1, true)) * bg_color_;
  out.depth = (weights * z_vals).sum(-1);
  out.weights = weights;
  return out;
}

torch::Tensor sample_pdf(const torch::Tensor &bins, const torch::Tensor &weights,
                         int n_samples, bool deterministic) {
  const int64_t R = weights.size(0), M = weights.size(1);

  // Normalise to a PDF, then a CDF with a leading zero to match the M + 1 edges.
  auto w = weights + 1e-5f;
  auto cdf = torch::cumsum(w / w.sum(-1, true), -1);
  cdf = torch::cat({torch::zeros({R, 1}, cdf.options()), cdf}, -1);

  auto u = deterministic ? torch::linspace(0.0f, 1.0f, n_samples, cdf.options())
                               .expand({R, n_samples})
                               .contiguous()
                         : torch::rand({R, n_samples}, cdf.options());

  // Locate the CDF bin each u falls in, then invert it linearly.
  auto inds = torch::searchsorted(cdf, u, /*out_int32=*/false, /*right=*/true);
  auto pair = torch::stack({(inds - 1).clamp_min(0), inds.clamp_max(M)}, -1);
  auto gather_pair = [&](const torch::Tensor &src) {
    return torch::gather(src.unsqueeze(1).expand({R, n_samples, M + 1}), 2, pair);
  };
  auto cdf_g = gather_pair(cdf);
  auto bins_g = gather_pair(bins);

  auto denom = cdf_g.index({"...", 1}) - cdf_g.index({"...", 0});
  denom = torch::where(denom < 1e-5f, torch::ones_like(denom), denom);
  auto t = (u - cdf_g.index({"...", 0})) / denom;
  return bins_g.index({"...", 0}) +
         t * (bins_g.index({"...", 1}) - bins_g.index({"...", 0}));
}

namespace {

// Per-ray linear interpolation of fp (defined at sorted xp) at query points.
torch::Tensor interp_per_ray(const torch::Tensor &query, const torch::Tensor &xp,
                             const torch::Tensor &fp) {
  auto inds = torch::searchsorted(xp, query, /*out_int32=*/false, /*right=*/true)
                  .clamp(1, xp.size(1) - 1);
  auto lo = inds - 1;
  auto x0 = torch::gather(xp, 1, lo), x1 = torch::gather(xp, 1, inds);
  auto y0 = torch::gather(fp, 1, lo), y1 = torch::gather(fp, 1, inds);
  return y0 + ((query - x0) / (x1 - x0 + 1e-7f)).clamp(0.0f, 1.0f) * (y1 - y0);
}

// Sample positions [R, S] -> S + 1 bin edges.
torch::Tensor to_edges(const torch::Tensor &z) {
  return torch::cat({head(z), midpoints(z), z.index({"...", Slice(-1, None)})}, -1);
}

}  // namespace

torch::Tensor interlevel_loss(const torch::Tensor &z_main, const torch::Tensor &w_main,
                              const torch::Tensor &z_prop, const torch::Tensor &w_prop) {
  // Compare normalised distributions; the main one is a fixed target.
  auto wm = (w_main / (w_main.sum(-1, true) + 1e-6)).detach();
  auto wp = w_prop / (w_prop.sum(-1, true) + 1e-6);

  // Proposal mass over each main bin is the difference of its interpolated CDF at
  // the bin edges. Penalise where the proposal underestimates it.
  auto cwp = torch::cat({torch::zeros({wp.size(0), 1}, wp.options()),
                         torch::cumsum(wp, -1)},
                        -1);
  auto cw_at = interp_per_ray(to_edges(z_main), to_edges(z_prop), cwp);
  auto excess = torch::relu(wm - (drop_first(cw_at) - drop_last(cw_at)));
  return (excess * excess / (wm + 1e-6)).sum(-1).mean();
}
