// viewer-tensor.cpp: tensor data analysis
//
// All functions operate on mmap'd files. The dequant dispatch reads raw
// blocks from the mapped data and converts to float via ggml's to_float.

#include "viewer-tensor.h"
#include "viewer-io.h"

#include <algorithm>
#include <cmath>
#include <cstring>
#include <limits>
#include <random>

static constexpr size_t CHUNK_ELEMENTS = 1u << 18;
static constexpr size_t RESERVOIR_SIZE = 1u << 14;

tensor_layout compute_layout(const tensor_entry &tensor) {
  tensor_layout layout;
  layout.width = 1;
  layout.height = 1;
  layout.depth = 1;

  int64_t nel = tensor.n_elements;
  if (nel == 0)
    return layout;

  int ndim = tensor.n_dims;

  if (ndim == 0) {
    layout.width = static_cast<size_t>(nel);
  } else if (ndim == 1) {
    size_t total = static_cast<size_t>(nel);
    size_t w =
        static_cast<size_t>(std::ceil(std::sqrt(static_cast<double>(total))));
    if (w == 0)
      w = 1;
    if (w > 1024)
      w = 1024;
    size_t h = (total + w - 1) / w;
    if (h == 0)
      h = 1;
    layout.width = w;
    layout.height = h;
  } else if (ndim == 2) {
    layout.width = static_cast<size_t>(std::max<int64_t>(1, tensor.ne[0]));
    layout.height = static_cast<size_t>(std::max<int64_t>(1, tensor.ne[1]));
  } else {
    layout.depth =
        static_cast<size_t>(std::max<int64_t>(1, tensor.ne[ndim - 1]));
    layout.width =
        static_cast<size_t>(std::max<int64_t>(1, tensor.ne[ndim - 2]));

    size_t plane = 1;
    for (int i = 0; i + 1 < ndim; ++i) {
      plane *= static_cast<size_t>(std::max<int64_t>(1, tensor.ne[i]));
    }
    layout.height = plane / layout.width;
    if (layout.height == 0)
      layout.height = 1;
  }

  return layout;
}

// read a contiguous range of elements from tensor data, converting to float.
// returns the number of floats written to dst.
static size_t read_elements(const mapped_file &file, const tensor_entry &tensor,
                            size_t element_offset, size_t count, float *dst) {
  size_t blk = tensor.block_size;
  size_t tsz = tensor.type_size;
  if (blk == 0 || tsz == 0)
    return 0;

  // clamp to tensor bounds
  if (element_offset >= static_cast<size_t>(tensor.n_elements))
    return 0;
  count = std::min(count,
                   static_cast<size_t>(tensor.n_elements) - element_offset);
  if (count == 0)
    return 0;

  // compute byte range
  size_t start_block = element_offset / blk;
  size_t end_block = (element_offset + count + blk - 1) / blk;
  size_t byte_end = tensor.file_offset + end_block * tsz;

  if (byte_end > file.size)
    return 0;

  // dequant block by block
  size_t written = 0;
  size_t skip = element_offset - start_block * blk;

  for (size_t b = start_block; b < end_block && written < count; ++b) {
    const uint8_t *block_ptr = file.data + tensor.file_offset + b * tsz;

    float tmp[256];
    size_t got = dequant_to_float(tensor.type, block_ptr, tmp, blk);
    if (got == 0)
      return written;

    // copy relevant elements
    for (size_t i = skip; i < got && written < count; ++i) {
      dst[written++] = tmp[i];
    }
    skip = 0;
  }

  return written;
}

bool tensor_read_tile(const mapped_file &file, const tensor_entry &tensor,
                      size_t slice, size_t x, size_t y, size_t width,
                      size_t height, tensor_tile &out, std::string &error) {
  out = {};

  const tensor_layout &layout = tensor.layout;

  if (x >= layout.width || y >= layout.height) {
    error = "tile outside tensor bounds";
    return false;
  }

  width = std::min(width, layout.width - x);
  height = std::min(height, layout.height - y);
  if (slice >= layout.depth)
    slice = layout.depth - 1;

  out.x = x;
  out.y = y;
  out.slice = slice;
  out.width = width;
  out.height = height;
  out.values.assign(width * height, 0.0f);
  out.mask.assign(width * height, 0);
  out.valid = 0;

  size_t slice_size = layout.width * layout.height;
  size_t base_offset = slice * slice_size;

  for (size_t row = 0; row < height; ++row) {
    size_t elem_start = base_offset + (y + row) * layout.width + x;
    if (elem_start >= static_cast<size_t>(tensor.n_elements))
      break;

    size_t row_count = std::min(
        width,
        static_cast<size_t>(tensor.n_elements) - elem_start);

    std::vector<float> row_buf(row_count);
    size_t got =
        read_elements(file, tensor, elem_start, row_count, row_buf.data());

    for (size_t col = 0; col < got; ++col) {
      size_t dst_idx = row * width + col;
      float val = row_buf[col];
      out.values[dst_idx] = val;
      out.mask[dst_idx] = 1;
      if (out.valid == 0) {
        out.min = val;
        out.max = val;
      } else {
        if (val < out.min)
          out.min = val;
        if (val > out.max)
          out.max = val;
      }
      ++out.valid;
    }
  }

  return true;
}

bool tensor_slice_stats(const mapped_file &file, const tensor_entry &tensor,
                        size_t slice, slice_stats &out, std::string &error) {
  out = {};

  const tensor_layout &layout = tensor.layout;
  if (slice >= layout.depth)
    slice = layout.depth - 1;

  size_t slice_size = layout.width * layout.height;
  size_t base_offset = slice * slice_size;
  if (base_offset >= static_cast<size_t>(tensor.n_elements))
    return true;

  size_t avail = std::min(
      slice_size,
      static_cast<size_t>(tensor.n_elements) - base_offset);
  if (avail == 0)
    return true;

  // stream through the slice, compute min/max + reservoir sample for
  // percentiles
  float min_val = std::numeric_limits<float>::infinity();
  float max_val = -std::numeric_limits<float>::infinity();
  size_t valid_count = 0;

  std::vector<float> reservoir;
  reservoir.reserve(RESERVOIR_SIZE);
  std::mt19937 rng(42);

  size_t remaining = avail;
  size_t offset = base_offset;
  std::vector<float> chunk(CHUNK_ELEMENTS);

  while (remaining > 0) {
    size_t n = std::min(remaining, CHUNK_ELEMENTS);
    size_t got = read_elements(file, tensor, offset, n, chunk.data());
    if (got == 0)
      break;

    for (size_t i = 0; i < got; ++i) {
      float v = chunk[i];
      if (!std::isfinite(v))
        continue;

      if (v < min_val)
        min_val = v;
      if (v > max_val)
        max_val = v;

      if (valid_count < RESERVOIR_SIZE) {
        reservoir.push_back(v);
      } else {
        std::uniform_int_distribution<size_t> dist(0, valid_count);
        size_t j = dist(rng);
        if (j < RESERVOIR_SIZE) {
          reservoir[j] = v;
        }
      }
      ++valid_count;
    }

    offset += got;
    remaining -= got;
  }

  out.computed = true;
  out.valid = valid_count;

  if (valid_count == 0)
    return true;

  // compute percentiles from reservoir
  std::sort(reservoir.begin(), reservoir.end());
  size_t n = reservoir.size();
  double max_idx = static_cast<double>(n - 1);

  auto interpolate = [&](double pct) -> float {
    double pos = pct * max_idx;
    size_t lo = static_cast<size_t>(std::floor(pos));
    size_t hi = std::min(lo + 1, n - 1);
    float w = static_cast<float>(pos - static_cast<double>(lo));
    return reservoir[lo] + w * (reservoir[hi] - reservoir[lo]);
  };

  float lower = interpolate(0.01);
  float upper = interpolate(0.99);

  if (!std::isfinite(lower) || lower < min_val)
    lower = min_val;
  if (!std::isfinite(upper) || upper > max_val)
    upper = max_val;
  if (lower > upper)
    std::swap(lower, upper);

  out.min = lower;
  out.max = upper;
  out.p_lower = 1.0f;
  out.p_upper = 99.0f;

  return true;
}

bool tensor_slice_histogram(const mapped_file &file, const tensor_entry &tensor,
                            size_t slice, size_t bin_count,
                            tensor_histogram &out, std::string &error) {
  out = {};
  if (bin_count == 0)
    return true;

  // get stats for range
  slice_stats stats;
  if (!tensor_slice_stats(file, tensor, slice, stats, error))
    return false;
  if (!stats.computed || stats.valid == 0) {
    out.bins.assign(bin_count, 0);
    return true;
  }

  float range_min = stats.min;
  float range_max = stats.max;
  float span = range_max - range_min;
  bool has_range =
      std::isfinite(range_min) && std::isfinite(range_max) && span > 0.0f;

  out.slice = slice;
  out.range_min = range_min;
  out.range_max = range_max;
  out.bins.assign(bin_count, 0);

  // stream and bin
  const tensor_layout &layout = tensor.layout;
  if (slice >= layout.depth)
    slice = layout.depth - 1;

  size_t slice_size = layout.width * layout.height;
  size_t base_offset = slice * slice_size;
  if (base_offset >= static_cast<size_t>(tensor.n_elements))
    return true;

  size_t avail = std::min(
      slice_size,
      static_cast<size_t>(tensor.n_elements) - base_offset);
  size_t remaining = avail;
  size_t offset = base_offset;
  std::vector<float> chunk(CHUNK_ELEMENTS);

  while (remaining > 0) {
    size_t n = std::min(remaining, CHUNK_ELEMENTS);
    size_t got = read_elements(file, tensor, offset, n, chunk.data());
    if (got == 0)
      break;

    for (size_t i = 0; i < got; ++i) {
      float v = chunk[i];
      if (!std::isfinite(v))
        continue;

      if (has_range) {
        if (v < range_min) {
          out.clipped_lo++;
          continue;
        }
        if (v > range_max) {
          out.clipped_hi++;
          continue;
        }
      }

      if (has_range && v == 0.0f) {
        out.zero_count++;
        continue;
      }

      size_t bin = 0;
      if (has_range) {
        float norm = (v - range_min) / span;
        if (norm <= 0.0f)
          bin = 0;
        else if (norm >= 1.0f)
          bin = bin_count - 1;
        else
          bin = static_cast<size_t>(norm * static_cast<float>(bin_count));
        if (bin >= bin_count)
          bin = bin_count - 1;
      }

      out.bins[bin]++;
      if (out.bins[bin] > out.max_bin)
        out.max_bin = out.bins[bin];
      out.total++;
    }

    offset += got;
    remaining -= got;
  }

  return true;
}

bool tensor_element_at(const mapped_file &file, const tensor_entry &tensor,
                       size_t slice, size_t x, size_t y, element_details &out,
                       std::string &error) {
  out = {};

  const tensor_layout &layout = tensor.layout;
  if (slice >= layout.depth)
    slice = layout.depth - 1;
  if (x >= layout.width)
    x = layout.width - 1;
  if (y >= layout.height)
    y = layout.height - 1;

  size_t slice_size = layout.width * layout.height;
  size_t element_index = slice * slice_size + y * layout.width + x;

  if (element_index >= static_cast<size_t>(tensor.n_elements)) {
    error = "element outside tensor";
    return false;
  }

  size_t blk = tensor.block_size;
  size_t tsz = tensor.type_size;
  if (blk == 0 || tsz == 0) {
    error = "unsupported tensor type";
    return false;
  }

  size_t block_index = element_index / blk;
  size_t index_in_block = element_index % blk;
  size_t block_byte_offset = block_index * tsz;
  size_t abs_offset = tensor.file_offset + block_byte_offset;

  if (abs_offset + tsz > file.size) {
    error = "tensor data outside file";
    return false;
  }

  // dequant the block
  float tmp[256];
  size_t got = dequant_to_float(tensor.type, file.data + abs_offset, tmp, blk);
  if (got == 0 || index_in_block >= got) {
    error = "dequant failed";
    return false;
  }

  out.element_index = element_index;
  out.element_count = static_cast<size_t>(tensor.n_elements);
  out.block_index = block_index;
  out.index_in_block = index_in_block;
  out.tensor_byte_offset = tensor.offset + block_byte_offset;
  out.file_byte_offset = abs_offset;
  out.value = tmp[index_in_block];
  out.valid = true;

  return true;
}
