stringzilla 5.0.2

Search, hash, sort, fingerprint, and fuzzy-match strings faster via SWAR, SIMD, and GPGPU
Documentation
/**
 *  @file c/stringzillas/levenshtein.cuh
 *  @brief Parallel Levenshtein & UTF-8 Levenshtein distances shim (CPU + CUDA backends).
 *  @author Ash Vardanian
 *  @date March 23, 2025
 */
#ifndef STRINGZILLAS_SZS_LEVENSHTEIN_CUH_
#define STRINGZILLAS_SZS_LEVENSHTEIN_CUH_
#include "stringzillas.cuh"

/**
 *  @brief Allocates a `levenshtein_backends_t` holding the `engine_type_` arm built from @p ctor_args, publishes the
 *         opaque handle, and folds the bad-alloc / success status reporting that every capability arm repeats.
 *  @tparam engine_type_ The concrete backend variant alternative to emplace (the only argument that varies per arm).
 */
template <typename engine_type_, typename... ctor_args_types_>
sz_status_t emplace_levenshtein_engine(szs_levenshtein_distances_t *engine_punned, char const **error_message,
                                       ctor_args_types_ &&...ctor_args) noexcept {
    auto engine = new (std::nothrow) levenshtein_backends_t(std::in_place_type_t<engine_type_>(),
                                                            engine_type_(std::forward<ctor_args_types_>(ctor_args)...));
    if (!engine)
        return propagate_error(sz::status_t::bad_alloc_k, error_message, "Failed to allocate Levenshtein engine");
    *engine_punned = reinterpret_cast<szs_levenshtein_distances_t>(engine);
    return propagate_error(sz::status_t::success_k, error_message);
}

/** @brief UTF-8 sibling of `emplace_levenshtein_engine` targeting `levenshtein_utf8_backends_t`. */
template <typename engine_type_, typename... ctor_args_types_>
sz_status_t emplace_levenshtein_utf8_engine(szs_levenshtein_distances_utf8_t *engine_punned, char const **error_message,
                                            ctor_args_types_ &&...ctor_args) noexcept {
    auto engine = new (std::nothrow) levenshtein_utf8_backends_t(
        std::in_place_type_t<engine_type_>(), engine_type_(std::forward<ctor_args_types_>(ctor_args)...));
    if (!engine)
        return propagate_error(sz::status_t::bad_alloc_k, error_message, "Failed to allocate UTF-8 Levenshtein engine");
    *engine_punned = reinterpret_cast<szs_levenshtein_distances_utf8_t>(engine);
    return propagate_error(sz::status_t::success_k, error_message);
}

/**
 *  @brief Cross-product (or symmetric self-similarity) dispatch shared by the byte-level and UTF-8 Levenshtein shims.
 *
 *  Builds an `szs::strided_rows<sz_size_t>` view over the caller's @p results matrix and invokes the migrated C++
 *  engine: the two-set overload `engine(queries, candidates, matrix, ...)` when @p candidates_container is non-null,
 *  or the symmetric overload `engine(queries, matrix, ...)` when it is null. The unsigned distance value type is
 *  `sz_size_t`, and @p results_row_stride counts elements (not bytes) between consecutive query rows.
 */
template <typename backends_type_, typename queries_type_, typename candidates_type_>
sz_status_t szs_levenshtein_cross_(                                                       //
    backends_type_ *engine, szs_device_scope_t device_punned,                             //
    queries_type_ const &queries_container, candidates_type_ const *candidates_container, //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(device_punned != nullptr && "Device must be initialized");
    sz_assert_(results != nullptr && "Results must not be null");

    auto *device = reinterpret_cast<device_scope_t *>(device_punned);
    auto const queries_count = queries_container.size();
    auto const candidates_count = candidates_container != nullptr ? candidates_container->size() : queries_count;
    auto results_matrix = szs::strided_rows<sz_size_t> {results, queries_count, candidates_count, results_row_stride};

    sz_status_t result = sz_success_k;
    auto variant_logic = [&](auto &engine_variant) {
        using engine_variant_t = std::decay_t<decltype(engine_variant)>;
        constexpr sz_capability_t engine_capability_k = engine_variant_t::capability_k;

        // GPU backends are only compatible with GPU scopes
        if constexpr (is_gpu_capability(engine_capability_k)) {
#if SZ_USE_CUDA
            if (std::holds_alternative<gpu_scope_t>(device->variants)) {
                auto &device_scope = std::get<gpu_scope_t>(device->variants);
                szs::cuda_status_t status = candidates_container != nullptr
                                                ? engine_variant(queries_container, *candidates_container,
                                                                 results_matrix, get_executor(device_scope),
                                                                 get_specs(device_scope))
                                                : engine_variant(queries_container, results_matrix, //
                                                                 get_executor(device_scope), get_specs(device_scope));
                result = propagate_error(status, error_message);
            }
            // Try ephemeral GPU on default scope (device 0)
            else if (std::holds_alternative<default_scope_t>(device->variants)) {
                auto &ctx = default_gpu_context();
                szs::cuda_status_t status = ctx.status != sz::status_t::success_k ? ctx.status
                                            : candidates_container != nullptr
                                                ? engine_variant(queries_container, *candidates_container,
                                                                 results_matrix, ctx.executor, ctx.specs)
                                                : engine_variant(queries_container, results_matrix, ctx.executor,
                                                                 ctx.specs);
                result = propagate_error(status, error_message);
            }
            else { result = propagate_error(sz::status_t::device_code_mismatch_k, error_message); }
#else
            result = propagate_error(sz::status_t::missing_gpu_k, error_message);
#endif // SZ_USE_CUDA
        }
        // CPU backends are only compatible with CPU scopes
        else {
            if (std::holds_alternative<default_scope_t>(device->variants)) {
                auto &device_scope = std::get<default_scope_t>(device->variants);
                sz::status_t status = candidates_container != nullptr
                                          ? engine_variant(queries_container, *candidates_container, results_matrix,
                                                           get_executor(device_scope), get_specs(device_scope))
                                          : engine_variant(queries_container, results_matrix, //
                                                           get_executor(device_scope), get_specs(device_scope));
                result = propagate_error(status, error_message);
            }
            else if (std::holds_alternative<cpu_scope_t>(device->variants)) {
                auto &device_scope = std::get<cpu_scope_t>(device->variants);
                sz::status_t status = candidates_container != nullptr
                                          ? engine_variant(queries_container, *candidates_container, results_matrix,
                                                           get_executor(device_scope), get_specs(device_scope))
                                          : engine_variant(queries_container, results_matrix, //
                                                           get_executor(device_scope), get_specs(device_scope));
                result = propagate_error(status, error_message);
            }
            else { result = propagate_error(sz::status_t::device_code_mismatch_k, error_message); }
        }
    };

    std::visit(variant_logic, engine->variants);
    return result;
}

extern "C" {

#pragma region Levenshtein Distances

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_init(                                         //
    sz_error_cost_t match, sz_error_cost_t mismatch, sz_error_cost_t open, sz_error_cost_t extend, //
    sz_memory_allocator_t const *alloc, sz_capability_t capabilities,                              //
    szs_levenshtein_distances_t *engine_punned, char const **error_message) {

    sz_unused_(alloc);        // Custom allocator not yet implemented, using default
    sz_unused_(capabilities); // Optional backends may be compiled out
    sz_assert_(engine_punned != nullptr && *engine_punned == nullptr && "Engine must be uninitialized");

    // If the gap opening and extension costs are identical we can use less memory
    auto const can_use_linear_costs = open == extend;
    auto const substitution_costs = szs::uniform_substitution_costs_t {match, mismatch};
    auto const linear_costs = szs::linear_gap_costs_t {open};
    auto const affine_costs = szs::affine_gap_costs_t {open, extend};

#if SZ_USE_ICELAKE
    bool const can_use_icelake = (capabilities & sz_cap_icelake_k) == sz_cap_icelake_k;
    if (can_use_icelake && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_icelake_t>(engine_punned, error_message, substitution_costs,
                                                                      linear_costs);
    else if (can_use_icelake)
        return emplace_levenshtein_engine<szs::affine_levenshtein_icelake_t>(engine_punned, error_message,
                                                                             substitution_costs, affine_costs);
#endif // SZ_USE_ICELAKE

#if SZ_USE_HASWELL
    bool const can_use_haswell = (capabilities & sz_cap_haswell_k) == sz_cap_haswell_k;
    if (can_use_haswell && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_haswell_t>(engine_punned, error_message, substitution_costs,
                                                                      linear_costs);
    else if (can_use_haswell)
        return emplace_levenshtein_engine<szs::affine_levenshtein_haswell_t>(engine_punned, error_message,
                                                                             substitution_costs, affine_costs);
#endif // SZ_USE_HASWELL

#if SZ_USE_NEON
    bool const can_use_neon = (capabilities & sz_cap_neon_k) == sz_cap_neon_k;
    if (can_use_neon && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_neon_t>(engine_punned, error_message, substitution_costs,
                                                                   linear_costs);
    else if (can_use_neon)
        return emplace_levenshtein_engine<szs::affine_levenshtein_neon_t>(engine_punned, error_message,
                                                                          substitution_costs, affine_costs);
#endif // SZ_USE_NEON

#if SZ_USE_RVV
    bool const can_use_rvv = (capabilities & sz_cap_rvv_k) == sz_cap_rvv_k;
    if (can_use_rvv && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_rvv_t>(engine_punned, error_message, substitution_costs,
                                                                  linear_costs);
    else if (can_use_rvv)
        return emplace_levenshtein_engine<szs::affine_levenshtein_rvv_t>(engine_punned, error_message,
                                                                         substitution_costs, affine_costs);
#endif // SZ_USE_RVV

    // GPU tiers are tested most-specific-first: a Hopper device reports the Kepler & base-CUDA bits too, so checking
    // base CUDA first would shadow the Hopper/Kepler engines. Hopper → Kepler → CUDA keeps each device on its best tier.
#if SZ_USE_HOPPER
    bool const can_use_hopper = (capabilities & sz_caps_ckh_k) == sz_caps_ckh_k;
    if (can_use_hopper && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_hopper_t>(engine_punned, error_message, substitution_costs,
                                                                     linear_costs);
    else if (can_use_hopper)
        return emplace_levenshtein_engine<szs::affine_levenshtein_hopper_t>(engine_punned, error_message,
                                                                            substitution_costs, affine_costs);
#endif // SZ_USE_HOPPER

#if SZ_USE_KEPLER
    bool const can_use_kepler = (capabilities & sz_caps_ck_k) == sz_caps_ck_k;
    if (can_use_kepler && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_kepler_t>(engine_punned, error_message, substitution_costs,
                                                                     linear_costs);
    else if (can_use_kepler)
        return emplace_levenshtein_engine<szs::affine_levenshtein_kepler_t>(engine_punned, error_message,
                                                                            substitution_costs, affine_costs);
#endif // SZ_USE_KEPLER

#if SZ_USE_CUDA
    bool const can_use_cuda = (capabilities & sz_cap_cuda_k) == sz_cap_cuda_k;
    if (can_use_cuda && can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_cuda_t>(engine_punned, error_message, substitution_costs,
                                                                   linear_costs);
    else if (can_use_cuda)
        return emplace_levenshtein_engine<szs::affine_levenshtein_cuda_t>(engine_punned, error_message,
                                                                          substitution_costs, affine_costs);
#endif // SZ_USE_CUDA

    if (can_use_linear_costs)
        return emplace_levenshtein_engine<szs::levenshtein_serial_t>(engine_punned, error_message, substitution_costs,
                                                                     linear_costs);
    else
        return emplace_levenshtein_engine<szs::affine_levenshtein_serial_t>(engine_punned, error_message,
                                                                            substitution_costs, affine_costs);
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances(                            //
    szs_levenshtein_distances_t engine_punned, szs_device_scope_t device_punned, //
    sz_sequence_t const *queries, sz_sequence_t const *candidates,               //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_u32tape(                      //
    szs_levenshtein_distances_t engine_punned, szs_device_scope_t device_punned,   //
    sz_sequence_u32tape_t const *queries, sz_sequence_u32tape_t const *candidates, //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_u32tape_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_u32tape_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_u64tape(                      //
    szs_levenshtein_distances_t engine_punned, szs_device_scope_t device_punned,   //
    sz_sequence_u64tape_t const *queries, sz_sequence_u64tape_t const *candidates, //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_u64tape_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_u64tape_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME void szs_levenshtein_distances_free(szs_levenshtein_distances_t engine_punned) {
    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    auto *engine = reinterpret_cast<levenshtein_backends_t *>(engine_punned);
    delete engine;
}

#pragma endregion Levenshtein Distances

#pragma region Levenshtein UTF8 Distances

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_utf8_init(                                    //
    sz_error_cost_t match, sz_error_cost_t mismatch, sz_error_cost_t open, sz_error_cost_t extend, //
    sz_memory_allocator_t const *alloc, sz_capability_t capabilities,                              //
    szs_levenshtein_distances_utf8_t *engine_punned, char const **error_message) {

    sz_unused_(alloc); // Custom allocator not yet implemented, using default
    sz_assert_(engine_punned != nullptr && *engine_punned == nullptr && "Engine must be uninitialized");

    // If the gap opening and extension costs are identical we can use less memory
    auto const can_use_linear_costs = open == extend;
    auto const substitution_costs = szs::uniform_substitution_costs_t {match, mismatch};
    auto const linear_costs = szs::linear_gap_costs_t {open};
    auto const affine_costs = szs::affine_gap_costs_t {open, extend};

#if SZ_USE_ICELAKE
    bool const can_use_icelake = (capabilities & sz_cap_icelake_k) != 0;
    if (can_use_icelake && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_icelake_t>(engine_punned, error_message,
                                                                                substitution_costs, linear_costs);
#endif // SZ_USE_ICELAKE

#if SZ_USE_HASWELL
    bool const can_use_haswell = (capabilities & sz_cap_haswell_k) != 0;
    if (can_use_haswell && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_haswell_t>(engine_punned, error_message,
                                                                                substitution_costs, linear_costs);
#endif // SZ_USE_HASWELL

#if SZ_USE_NEON
    bool const can_use_neon = (capabilities & sz_cap_neon_k) != 0;
    if (can_use_neon && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_neon_t>(engine_punned, error_message,
                                                                             substitution_costs, linear_costs);
#endif // SZ_USE_NEON

#if SZ_USE_RVV
    bool const can_use_rvv = (capabilities & sz_cap_rvv_k) != 0;
    if (can_use_rvv && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_rvv_t>(engine_punned, error_message,
                                                                            substitution_costs, linear_costs);
#endif // SZ_USE_RVV

#if SZ_USE_CUDA
    bool const can_use_cuda = (capabilities & sz_cap_cuda_k) == sz_cap_cuda_k;
    if (can_use_cuda && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_cuda_t>(engine_punned, error_message,
                                                                             substitution_costs, linear_costs);
#endif // SZ_USE_CUDA

    bool const can_use_serial = (capabilities & sz_cap_serial_k) == sz_cap_serial_k;
    if (can_use_serial && can_use_linear_costs)
        return emplace_levenshtein_utf8_engine<szs::levenshtein_utf8_serial_t>(engine_punned, error_message,
                                                                               substitution_costs, linear_costs);
    else
        return emplace_levenshtein_utf8_engine<szs::affine_levenshtein_utf8_serial_t>(engine_punned, error_message,
                                                                                      substitution_costs, affine_costs);

    return propagate_error(sz::status_t::unknown_k, error_message, "No supported UTF-8 Levenshtein backends available");
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_utf8(                            //
    szs_levenshtein_distances_utf8_t engine_punned, szs_device_scope_t device_punned, //
    sz_sequence_t const *queries, sz_sequence_t const *candidates,                    //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_utf8_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_utf8_u32tape(                    //
    szs_levenshtein_distances_utf8_t engine_punned, szs_device_scope_t device_punned, //
    sz_sequence_u32tape_t const *queries, sz_sequence_u32tape_t const *candidates,    //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_utf8_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_u32tape_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_u32tape_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME sz_status_t szs_levenshtein_distances_utf8_u64tape(                    //
    szs_levenshtein_distances_utf8_t engine_punned, szs_device_scope_t device_punned, //
    sz_sequence_u64tape_t const *queries, sz_sequence_u64tape_t const *candidates,    //
    sz_size_t *results, sz_size_t results_row_stride, char const **error_message) {

    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    sz_assert_(queries != nullptr && "Query texts cannot be null");
    auto *engine = reinterpret_cast<levenshtein_utf8_backends_t *>(engine_punned);
    auto queries_container = sz_sequence_u64tape_as_cpp_container_t {queries};
    auto candidates_container = sz_sequence_u64tape_as_cpp_container_t {candidates};
    return szs_levenshtein_cross_(                                                                         //
        engine, device_punned, queries_container, candidates != nullptr ? &candidates_container : nullptr, //
        results, results_row_stride, error_message);
}

SZ_API_RUNTIME void szs_levenshtein_distances_utf8_free(szs_levenshtein_distances_utf8_t engine_punned) {
    sz_assert_(engine_punned != nullptr && "Engine must be initialized");
    auto *engine = reinterpret_cast<levenshtein_utf8_backends_t *>(engine_punned);
    delete engine;
}

#pragma endregion Levenshtein UTF8 Distances
}

#endif // STRINGZILLAS_SZS_LEVENSHTEIN_CUH_