#include <algorithm>
#include <charconv>
#include <chrono>
#include <cstdint>
#include <cstdlib>
#include <iomanip>
#include <iostream>
#include <limits>
#include <numeric>
#include <random>
#include <stdexcept>
#include <string>
#include <string_view>
#include <system_error>
#include <vector>

namespace {

struct Options {
    std::size_t width = 4096;
    std::size_t height = 2048;
    std::size_t tile = 32;
    std::size_t passes = 6;
    std::size_t rounds = 7;
};

struct TileQuery {
    std::size_t linear_x;
    std::size_t linear_y;
    std::size_t tiled_base;
};

struct TimedResult {
    double milliseconds;
    std::uint64_t checksum;
};

[[noreturn]] void fail(const std::string& message) {
    throw std::runtime_error(message);
}

std::size_t parse_positive(std::string_view text, std::string_view name) {
    std::size_t value = 0;
    const char* begin = text.data();
    const char* end = begin + text.size();
    const auto result = std::from_chars(begin, end, value);
    if (result.ec != std::errc{} || result.ptr != end || value == 0) {
        fail("--" + std::string(name) + " must be a positive integer");
    }
    return value;
}

Options parse_options(int argc, char** argv) {
    Options options;

    for (int index = 1; index < argc; ++index) {
        const std::string_view argument(argv[index]);
        if (argument == "--help" || argument == "-h") {
            std::cout
                << "Usage: tiling-locality-benchmark [options]\n\n"
                << "  --width=N    image width in pixels (default 4096)\n"
                << "  --height=N   image height in pixels (default 2048)\n"
                << "  --tile=N     square tile edge in pixels (default 32)\n"
                << "  --passes=N   full-image reads per sample (default 6)\n"
                << "  --rounds=N   timed samples per layout (default 7)\n";
            std::exit(0);
        }

        const auto equals = argument.find('=');
        if (equals == std::string_view::npos || argument.size() < 2 ||
            argument.substr(0, 2) != "--") {
            fail("unknown argument: " + std::string(argument));
        }

        const std::string_view name = argument.substr(2, equals - 2);
        const std::size_t value = parse_positive(argument.substr(equals + 1), name);
        if (name == "width") {
            options.width = value;
        } else if (name == "height") {
            options.height = value;
        } else if (name == "tile") {
            options.tile = value;
        } else if (name == "passes") {
            options.passes = value;
        } else if (name == "rounds") {
            options.rounds = value;
        } else {
            fail("unknown option: --" + std::string(name));
        }
    }

    if (options.width % options.tile != 0 || options.height % options.tile != 0) {
        fail("width and height must both be multiples of tile");
    }
    if (options.width > std::numeric_limits<std::size_t>::max() / options.height) {
        fail("image dimensions overflow size_t");
    }
    return options;
}

std::vector<std::uint32_t> make_linear_image(std::size_t pixel_count) {
    std::vector<std::uint32_t> image(pixel_count);
    for (std::size_t index = 0; index < pixel_count; ++index) {
        const auto value = static_cast<std::uint32_t>(index);
        image[index] = (value * 2654435761u) >> 16;
    }
    return image;
}

std::vector<std::uint32_t> make_tiled_image(
    const std::vector<std::uint32_t>& linear,
    const Options& options)
{
    std::vector<std::uint32_t> tiled(linear.size());
    const std::size_t tiles_x = options.width / options.tile;
    const std::size_t tiles_y = options.height / options.tile;
    const std::size_t tile_area = options.tile * options.tile;

    for (std::size_t tile_y = 0; tile_y < tiles_y; ++tile_y) {
        for (std::size_t tile_x = 0; tile_x < tiles_x; ++tile_x) {
            const std::size_t tiled_base = (tile_y * tiles_x + tile_x) * tile_area;
            const std::size_t source_x = tile_x * options.tile;
            const std::size_t source_y = tile_y * options.tile;

            for (std::size_t local_y = 0; local_y < options.tile; ++local_y) {
                const std::size_t source_row =
                    (source_y + local_y) * options.width + source_x;
                const std::size_t tiled_row = tiled_base + local_y * options.tile;
                std::copy_n(
                    linear.data() + source_row,
                    options.tile,
                    tiled.data() + tiled_row);
            }
        }
    }
    return tiled;
}

std::vector<TileQuery> make_queries(const Options& options) {
    const std::size_t tiles_x = options.width / options.tile;
    const std::size_t tiles_y = options.height / options.tile;
    const std::size_t tile_area = options.tile * options.tile;
    std::vector<TileQuery> queries;
    queries.reserve(tiles_x * tiles_y);

    for (std::size_t tile_y = 0; tile_y < tiles_y; ++tile_y) {
        for (std::size_t tile_x = 0; tile_x < tiles_x; ++tile_x) {
            const std::size_t tile_index = tile_y * tiles_x + tile_x;
            queries.push_back({
                tile_x * options.tile,
                tile_y * options.tile,
                tile_index * tile_area,
            });
        }
    }

    std::mt19937 generator(0xC0FFEEu);
    std::shuffle(queries.begin(), queries.end(), generator);
    return queries;
}

std::uint64_t sum_linear_tiles(
    const std::vector<std::uint32_t>& image,
    const std::vector<TileQuery>& queries,
    const Options& options,
    std::size_t passes)
{
    std::uint64_t checksum = 0;
    for (std::size_t pass = 0; pass < passes; ++pass) {
        for (const TileQuery& query : queries) {
            for (std::size_t local_y = 0; local_y < options.tile; ++local_y) {
                const std::uint32_t* row = image.data()
                    + (query.linear_y + local_y) * options.width
                    + query.linear_x;
                for (std::size_t local_x = 0; local_x < options.tile; ++local_x) {
                    checksum += row[local_x];
                }
            }
        }
    }
    return checksum;
}

std::uint64_t sum_tiled_tiles(
    const std::vector<std::uint32_t>& image,
    const std::vector<TileQuery>& queries,
    const Options& options,
    std::size_t passes)
{
    std::uint64_t checksum = 0;
    for (std::size_t pass = 0; pass < passes; ++pass) {
        for (const TileQuery& query : queries) {
            for (std::size_t local_y = 0; local_y < options.tile; ++local_y) {
                const std::uint32_t* row = image.data()
                    + query.tiled_base
                    + local_y * options.tile;
                for (std::size_t local_x = 0; local_x < options.tile; ++local_x) {
                    checksum += row[local_x];
                }
            }
        }
    }
    return checksum;
}

template <typename Operation>
TimedResult measure(Operation&& operation) {
    const auto start = std::chrono::steady_clock::now();
    const std::uint64_t checksum = operation();
    const auto finish = std::chrono::steady_clock::now();
    const std::chrono::duration<double, std::milli> elapsed = finish - start;
    return {elapsed.count(), checksum};
}

double median(std::vector<double> values) {
    std::sort(values.begin(), values.end());
    const std::size_t middle = values.size() / 2;
    if (values.size() % 2 == 0) {
        return (values[middle - 1] + values[middle]) / 2.0;
    }
    return values[middle];
}

void print_result(
    std::string_view name,
    const std::vector<double>& samples,
    double bytes_read)
{
    const double middle = median(samples);
    const double gib_per_second = bytes_read / (middle / 1000.0) / (1024.0 * 1024.0 * 1024.0);
    std::cout << std::left << std::setw(7) << name << std::right
              << ": median " << std::fixed << std::setprecision(3) << middle << " ms, "
              << std::setprecision(2) << gib_per_second << " GiB/s\n";
}

} // namespace

int main(int argc, char** argv) {
    try {
        const Options options = parse_options(argc, argv);
        const std::size_t pixel_count = options.width * options.height;
        const double bytes_per_layout =
            static_cast<double>(pixel_count * sizeof(std::uint32_t));

        std::cout << "Image: " << options.width << 'x' << options.height
                  << ", tile: " << options.tile << 'x' << options.tile
                  << ", " << std::fixed << std::setprecision(1)
                  << bytes_per_layout / (1024.0 * 1024.0) << " MiB per layout\n"
                  << "Workload: visit every logical tile in a fixed random order, "
                  << options.passes << " pass(es) per sample\n";

        const std::vector<std::uint32_t> linear = make_linear_image(pixel_count);
        const std::vector<std::uint32_t> tiled = make_tiled_image(linear, options);
        const std::vector<TileQuery> queries = make_queries(options);

        const std::uint64_t warm_linear =
            sum_linear_tiles(linear, queries, options, 1);
        const std::uint64_t warm_tiled =
            sum_tiled_tiles(tiled, queries, options, 1);
        if (warm_linear != warm_tiled) {
            fail("layout conversion failed: warm-up checksums differ");
        }

        std::vector<double> linear_samples;
        std::vector<double> tiled_samples;
        linear_samples.reserve(options.rounds);
        tiled_samples.reserve(options.rounds);
        std::uint64_t linear_checksum = 0;
        std::uint64_t tiled_checksum = 0;

        for (std::size_t round = 0; round < options.rounds; ++round) {
            const auto run_linear = [&] {
                const TimedResult result = measure([&] {
                    return sum_linear_tiles(linear, queries, options, options.passes);
                });
                linear_samples.push_back(result.milliseconds);
                linear_checksum = result.checksum;
            };
            const auto run_tiled = [&] {
                const TimedResult result = measure([&] {
                    return sum_tiled_tiles(tiled, queries, options, options.passes);
                });
                tiled_samples.push_back(result.milliseconds);
                tiled_checksum = result.checksum;
            };

            if (round % 2 == 0) {
                run_linear();
                run_tiled();
            } else {
                run_tiled();
                run_linear();
            }
        }

        if (linear_checksum != tiled_checksum) {
            fail("measured checksums differ");
        }

        const double bytes_read = bytes_per_layout * static_cast<double>(options.passes);
        print_result("linear", linear_samples, bytes_read);
        print_result("tiled", tiled_samples, bytes_read);
        const double speedup = median(linear_samples) / median(tiled_samples);
        std::cout << "speedup: " << std::fixed << std::setprecision(2)
                  << speedup << "x\n"
                  << "checksum: " << linear_checksum << "\n"
                  << "Note: this models 2D locality on a CPU; it is not a direct "
                  << "Intel GPU hardware benchmark.\n";
        return 0;
    } catch (const std::exception& exception) {
        std::cerr << "error: " << exception.what() << '\n';
        return 1;
    }
}
