packages feed

eigen-hhlo-0.1.0.0: cbits/gpu/eigenhhlo_cusolver.cu

/* eigen-hhlo GPU custom-call library — cuSOLVER wrappers.
 *
 * Conforms to XLA GPU custom-call ABI (API version 3):
 *   void target(CUstream stream, void** buffers,
 *               const char* opaque, size_t opaque_len,
 *               XlaCustomCallStatus* status);
 *
 * Buffer layout: [in0, in1, ..., out0, out1, ...]
 *
 * Compile:
 *   cd cbits/gpu && bash build.sh
 */

#include <cuda.h>
#include <cuda_runtime.h>
#include <cusolverDn.h>
#include <stdlib.h>
#include <string.h>
#include <stdio.h>
#include <stdint.h>
#include <unordered_map>

/* --------------------------------------------------------------------------
 * XlaCustomCallStatus placeholder (opaque struct pointer)
 * -------------------------------------------------------------------------- */
struct XlaCustomCallStatus;

/* --------------------------------------------------------------------------
 * Opaque string parser: key=value,...
 * -------------------------------------------------------------------------- */
static int parse_int_kv(const char* opaque, const char* key, int default_val) {
    if (!opaque || !key) return default_val;
    size_t keylen = strlen(key);
    const char* p = opaque;
    while (*p) {
        if (strncmp(p, key, keylen) == 0 && p[keylen] == '=') {
            p += keylen + 1;
            return (int)strtol(p, NULL, 10);
        }
        while (*p && *p != ',') p++;
        if (*p == ',') p++;
    }
    return default_val;
}

static char parse_char_kv(const char* opaque, const char* key, char default_val) {
    if (!opaque || !key) return default_val;
    size_t keylen = strlen(key);
    const char* p = opaque;
    while (*p) {
        if (strncmp(p, key, keylen) == 0 && p[keylen] == '=') {
            p += keylen + 1;
            return *p;
        }
        while (*p && *p != ',') p++;
        if (*p == ',') p++;
    }
    return default_val;
}

/* --------------------------------------------------------------------------
 * cuSOLVER handle cache: one handle per CUDA stream
 * -------------------------------------------------------------------------- */
static thread_local std::unordered_map<CUstream, cusolverDnHandle_t> g_handleCache;

static cusolverDnHandle_t getHandle(CUstream stream) {
    auto it = g_handleCache.find(stream);
    if (it != g_handleCache.end()) return it->second;
    cusolverDnHandle_t h;
    cusolverDnCreate(&h);
    cusolverDnSetStream(h, (cudaStream_t)stream);
    g_handleCache[stream] = h;
    return h;
}

/* --------------------------------------------------------------------------
 * Helper: map char to cublasFillMode_t
 * -------------------------------------------------------------------------- */
static cublasFillMode_t charToFillMode(char uplo) {
    return (uplo == 'U' || uplo == 'u') ? CUBLAS_FILL_MODE_UPPER
                                        : CUBLAS_FILL_MODE_LOWER;
}

/* --------------------------------------------------------------------------
 * Helper: check cuSOLVER status and print errors
 * -------------------------------------------------------------------------- */
static void checkCusolver(cusolverStatus_t status, const char* func) {
    if (status != CUSOLVER_STATUS_SUCCESS) {
        fprintf(stderr, "eigen-hhlo GPU: %s failed with status %d\n", func, status);
    }
}

/* ==========================================================================
 * Cholesky: eigenhhlo_dpotrf
 *   buffers[0] = A (n×n, double, input)
 *   buffers[1] = L or U (n×n, double, output)
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dpotrf(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int n = parse_int_kv(opaque, "n", 1);
    char uplo_char = parse_char_kv(opaque, "uplo", 'L');
    cublasFillMode_t uplo = charToFillMode(uplo_char);

    const double* A = (const double*)buffers[0];
    double* outA = (double*)buffers[1];

    cudaMemcpyAsync(outA, A, n * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    int lwork;
    checkCusolver(cusolverDnDpotrf_bufferSize(handle, uplo, n, outA, n, &lwork),
                  "cusolverDnDpotrf_bufferSize");

    double* work;
    cudaMallocAsync(&work, lwork * sizeof(double), stream);
    int* devInfo;
    cudaMallocAsync(&devInfo, sizeof(int), stream);

    checkCusolver(cusolverDnDpotrf(handle, uplo, n, outA, n, work, lwork, devInfo),
                  "cusolverDnDpotrf");

    cudaFreeAsync(work, stream);
    cudaFreeAsync(devInfo, stream);
}

/* ==========================================================================
 * LU: eigenhhlo_dgetrf
 *   buffers[0] = A (m×n, double, input)
 *   buffers[1] = LU (m×n, double, output)
 *   buffers[2] = pivots (min(m,n), int32, output)
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dgetrf(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int m = parse_int_kv(opaque, "m", 1);
    int n = parse_int_kv(opaque, "n", 1);

    const double* A = (const double*)buffers[0];
    double* outLU = (double*)buffers[1];
    int32_t* outPiv = (int32_t*)buffers[2];

    cudaMemcpyAsync(outLU, A, m * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    int lwork;
    checkCusolver(cusolverDnDgetrf_bufferSize(handle, m, n, outLU, m, &lwork),
                  "cusolverDnDgetrf_bufferSize");

    double* work;
    cudaMallocAsync(&work, lwork * sizeof(double), stream);
    int* devInfo;
    cudaMallocAsync(&devInfo, sizeof(int), stream);

    checkCusolver(cusolverDnDgetrf(handle, m, n, outLU, m, work, (int*)outPiv, devInfo),
                  "cusolverDnDgetrf");

    cudaFreeAsync(work, stream);
    cudaFreeAsync(devInfo, stream);
}

/* ==========================================================================
 * QR factorization: eigenhhlo_dgeqrf
 *   buffers[0] = A (m×n, double, input)
 *   buffers[1] = A overwritten with R+reflectors (m×n, double, output)
 *   buffers[2] = tau (min(m,n), double, output)
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dgeqrf(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int m = parse_int_kv(opaque, "m", 1);
    int n = parse_int_kv(opaque, "n", 1);

    const double* A = (const double*)buffers[0];
    double* outA = (double*)buffers[1];
    double* outTau = (double*)buffers[2];

    cudaMemcpyAsync(outA, A, m * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    int lwork;
    checkCusolver(cusolverDnDgeqrf_bufferSize(handle, m, n, outA, m, &lwork),
                  "cusolverDnDgeqrf_bufferSize");

    double* work;
    cudaMallocAsync(&work, lwork * sizeof(double), stream);
    int* devInfo;
    cudaMallocAsync(&devInfo, sizeof(int), stream);

    checkCusolver(cusolverDnDgeqrf(handle, m, n, outA, m, outTau, work, lwork, devInfo),
                  "cusolverDnDgeqrf");

    cudaFreeAsync(work, stream);
    cudaFreeAsync(devInfo, stream);
}

/* ==========================================================================
 * Generate Q from QR: eigenhhlo_dorgqr
 *   buffers[0] = A (m×n, double, contains reflectors from dgeqrf)
 *   buffers[1] = tau (k, double)
 *   buffers[2] = Q (m×m, double, output)
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dorgqr(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int m = parse_int_kv(opaque, "m", 1);
    int n = parse_int_kv(opaque, "n", 1);
    int k = parse_int_kv(opaque, "k", 1);

    const double* A = (const double*)buffers[0];
    const double* tau = (const double*)buffers[1];
    double* outQ = (double*)buffers[2];

    /* Initialize Q to identity by zeroing and setting diagonal, then copy reflectors */
    cudaMemsetAsync(outQ, 0, m * m * sizeof(double), stream);
    /* Copy reflectors to the left part of Q */
    cudaMemcpyAsync(outQ, A, m * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    int lwork;
    checkCusolver(cusolverDnDorgqr_bufferSize(handle, m, m, k, outQ, m, tau, &lwork),
                  "cusolverDnDorgqr_bufferSize");

    double* work;
    cudaMallocAsync(&work, lwork * sizeof(double), stream);
    int* devInfo;
    cudaMallocAsync(&devInfo, sizeof(int), stream);

    checkCusolver(cusolverDnDorgqr(handle, m, m, k, outQ, m, tau, work, lwork, devInfo),
                  "cusolverDnDorgqr");

    cudaFreeAsync(work, stream);
    cudaFreeAsync(devInfo, stream);
}

/* ==========================================================================
 * Symmetric eigenvalue: eigenhhlo_dsyevd
 *   buffers[0] = A (n×n, double, symmetric, input)
 *   buffers[1] = eigenvalues (n, double, output)
 *   buffers[2] = eigenvectors (n×n, double, output)
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dsyevd(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int n = parse_int_kv(opaque, "n", 1);
    char uplo_char = parse_char_kv(opaque, "uplo", 'L');
    cublasFillMode_t uplo = charToFillMode(uplo_char);

    const double* A = (const double*)buffers[0];
    double* outW = (double*)buffers[1];
    double* outV = (double*)buffers[2];

    cudaMemcpyAsync(outV, A, n * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    int lwork;
    checkCusolver(cusolverDnDsyevd_bufferSize(handle, CUSOLVER_EIG_MODE_VECTOR, uplo, n, outV, n, outW, &lwork),
                  "cusolverDnDsyevd_bufferSize");

    double* work;
    cudaMallocAsync(&work, lwork * sizeof(double), stream);
    int* devInfo;
    cudaMallocAsync(&devInfo, sizeof(int), stream);

    checkCusolver(cusolverDnDsyevd(handle, CUSOLVER_EIG_MODE_VECTOR, uplo, n, outV, n, outW, work, lwork, devInfo),
                  "cusolverDnDsyevd");

    cudaFreeAsync(work, stream);
    cudaFreeAsync(devInfo, stream);
}

/* --------------------------------------------------------------------------
 * Transpose an m×n column-major matrix to an n×m column-major matrix.
 * -------------------------------------------------------------------------- */
__global__ void transpose_colmajor(double* out, const double* in, int m, int n)
{
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < m * n) {
        int i = idx % m;   // row in input
        int j = idx / m;   // col in input
        out[j + i * n] = in[i + j * m];
    }
}

/* ==========================================================================
 * SVD: eigenhhlo_dgesvd
 *   buffers[0] = A (m×n, double, input)
 *   buffers[1] = U (m×m or m×min, double, output)
 *   buffers[2] = S (min(m,n), double, output)
 *   buffers[3] = Vt (n×n or min×n, double, output)
 *
 * Workaround: cusolverDnDgesvd returns CUSOLVER_STATUS_INVALID_VALUE
 * when m < n on some driver/GPU combinations.  We transpose A, call
 * dgesvd on A^T (where n >= m), and swap the U / Vt output buffers.
 * ========================================================================== */
extern "C" __attribute__((visibility("default")))
void eigenhhlo_dgesvd(CUstream stream, void** buffers,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int m = parse_int_kv(opaque, "m", 1);
    int n = parse_int_kv(opaque, "n", 1);
    char jobu  = parse_char_kv(opaque, "jobu", 'A');
    char jobvt = parse_char_kv(opaque, "jobvt", 'A');

    const double* A = (const double*)buffers[0];
    double* outU  = (double*)buffers[1];
    double* outS  = (double*)buffers[2];
    double* outVt = (double*)buffers[3];

    /* cuSOLVER destroys A; allocate a temporary workspace copy */
    double* awork;
    cudaMallocAsync(&awork, m * n * sizeof(double), stream);
    cudaMemcpyAsync(awork, A, m * n * sizeof(double), cudaMemcpyDeviceToDevice, stream);

    cusolverDnHandle_t handle = getHandle(stream);

    if (m < n) {
        /* Transpose A to A^T (n×m) so that n >= m */
        double* temp;
        cudaMallocAsync(&temp, m * n * sizeof(double), stream);
        int threads = 256;
        int blocks = (m * n + threads - 1) / threads;
        transpose_colmajor<<<blocks, threads, 0, stream>>>(temp, awork, m, n);

        int lwork;
        checkCusolver(cusolverDnDgesvd_bufferSize(handle, n, m, &lwork),
                      "cusolverDnDgesvd_bufferSize");
        double* work;
        cudaMallocAsync(&work, lwork * sizeof(double), stream);
        int* devInfo;
        cudaMallocAsync(&devInfo, sizeof(int), stream);

        /* Swap U and VT outputs: dgesvd on A^T writes U' to what we call VT
         * and VT' to what we call U.  ldu=n (rows of U'), ldvt=m (rows of VT'). */
        checkCusolver(cusolverDnDgesvd(handle, jobvt, jobu, n, m, temp, n,
                                       outS, outVt, n, outU, m,
                                       work, lwork, NULL, devInfo),
                      "cusolverDnDgesvd");

        cudaFreeAsync(work, stream);
        cudaFreeAsync(devInfo, stream);
        cudaFreeAsync(temp, stream);
    } else {
        int k = (m < n) ? m : n;
        int ldvt = (jobvt == 'S' || jobvt == 's') ? k : n;

        int lwork;
        checkCusolver(cusolverDnDgesvd_bufferSize(handle, m, n, &lwork),
                      "cusolverDnDgesvd_bufferSize");
        double* work;
        cudaMallocAsync(&work, lwork * sizeof(double), stream);
        int* devInfo;
        cudaMallocAsync(&devInfo, sizeof(int), stream);

        checkCusolver(cusolverDnDgesvd(handle, jobu, jobvt, m, n, awork, m,
                                       outS, outU, m, outVt, ldvt,
                                       work, lwork, NULL, devInfo),
                      "cusolverDnDgesvd");

        cudaFreeAsync(work, stream);
        cudaFreeAsync(devInfo, stream);
    }

    cudaFreeAsync(awork, stream);
}