// Bit-preserving FP8 packing. Workers own complete destination words, including
// padding; no quantization, byte stores or cross-worker read/modify/write.
#include <poplar/Vertex.hpp>
using namespace poplar;

class REARRANGE_VERTEX_NAME : public MultiVertex {
public:
  Input<Vector<unsigned char, VectorLayout::ONE_PTR>> source;
  Output<Vector<unsigned char, VectorLayout::ONE_PTR>> destination;
  unsigned logicalRows, physicalRows, targetOrder, logicalColumns, physicalColumns, matrices;

  bool compute(unsigned worker) {
    const unsigned rows = REARRANGE_LOGICAL_ROWS ? REARRANGE_LOGICAL_ROWS : logicalRows;
    const unsigned pr = REARRANGE_PHYSICAL_ROWS ? REARRANGE_PHYSICAL_ROWS : physicalRows;
    const unsigned columns = REARRANGE_LOGICAL_COLUMNS ? REARRANGE_LOGICAL_COLUMNS : logicalColumns;
    const unsigned pc = REARRANGE_PHYSICAL_COLUMNS ? REARRANGE_PHYSICAL_COLUMNS : physicalColumns;
    constexpr unsigned inner = 32;
    constexpr unsigned rb = REARRANGE_ROW_BLOCK;
    constexpr unsigned cb = REARRANGE_COLUMN_BLOCK;
    for (unsigned matrix = 0; matrix < matrices; ++matrix) {
      const unsigned char *input = &source[0] + matrix * rows * columns;
      unsigned *output = reinterpret_cast<unsigned *>(&destination[0]);
      // Common aligned case is one word load. Only clipped/unaligned row
      // boundaries use scalar loads; zero-fill never reads past the input.
      auto load = [&](unsigned row, unsigned column) -> unsigned {
        if (row >= rows || column >= columns) return 0;
        const unsigned offset = row * columns + column;
        if (((matrix * rows * columns + offset) & 3) == 0 && column + 4 <= columns)
          return *reinterpret_cast<const unsigned *>(input + offset);
        unsigned word = 0;
        for (unsigned lane = 0; lane < 4 && column + lane < columns; ++lane)
          word |= unsigned(input[offset + lane]) << (lane * 8);
        return word;
      };
#if REARRANGE_TARGET_ORDER < 2
      for (unsigned row = worker; row < pr; row += 6) {
        for (unsigned column = 0; column < pc; column += 4) {
#if REARRANGE_TARGET_ORDER == 0
          const unsigned offset = (column / inner) * pr * inner * matrices
              + matrix * pr * inner + row * inner + column % inner;
#else
          const unsigned panel = (row / cb) * (pc / inner) + column / inner;
          const unsigned offset = matrix * pr * pc + panel * inner * cb
              + row % cb * inner + column % inner;
#endif
          output[offset / 4] = load(row, column);
        }
      }
#else
      // A 4x4 byte transpose turns four source word loads into four destination
      // word stores. Rows within each 32x16 AMP micro-panel are contiguous.
      for (unsigned row = worker * 4; row < pr; row += 24) {
        for (unsigned column = 0; column < pc; column += 4) {
          const unsigned a = load(row, column), b = load(row + 1, column);
          const unsigned c = load(row + 2, column), d = load(row + 3, column);
          const unsigned panel = (row / rb) * (pc / cb) * (rb / inner)
              + (column / cb) * (rb / inner) + row % rb / inner;
          const unsigned base = matrix * pr * pc + panel * inner * cb + row % inner;
          #pragma unroll
          for (unsigned lane = 0; lane < 4; ++lane) {
            const unsigned shift = lane * 8;
            const unsigned word = ((a >> shift) & 255) | (((b >> shift) & 255) << 8)
                | (((c >> shift) & 255) << 16) | (((d >> shift) & 255) << 24);
            output[(base + (column % cb + lane) * inner) / 4] = word;
          }
        }
      }
#endif
    }
    return true;
  }
};
