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);
}