packages feed

eigen-hhlo-0.1.0.0: cbits/cpu/eigenhhlo_lapack.c

/* eigen-hhlo CPU custom-call library — LAPACK wrappers.
 *
 * Conforms to XLA CPU custom-call ABI (API version 2):
 *   void target(void* out, void** in,
 *               const char* opaque, size_t opaque_len,
 *               XlaCustomCallStatus* status);
 *
 * For multi-result custom calls, `out` is a void** array of output buffers.
 *
 * Compile:
 *   cd cbits/cpu && bash build.sh
 */

#include <stdlib.h>
#include <string.h>
#include <stdio.h>
#include <stdint.h>

/* --------------------------------------------------------------------------
 * XlaCustomCallStatus placeholder (opaque struct pointer)
 * -------------------------------------------------------------------------- */
typedef struct XlaCustomCallStatus_ 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;
}

/* --------------------------------------------------------------------------
 * LAPACK declarations (Fortran ABI: all args by pointer, trailing underscore)
 * -------------------------------------------------------------------------- */
extern void dpotrf_(char* uplo, int* n, double* a, int* lda, int* info);
extern void dgetrf_(int* m, int* n, double* a, int* lda, int* ipiv, int* info);
extern void dgeqrf_(int* m, int* n, double* a, int* lda, double* tau,
                    double* work, int* lwork, int* info);
extern void dorgqr_(int* m, int* n, int* k, double* a, int* lda, double* tau,
                    double* work, int* lwork, int* info);
extern void dsyevd_(char* jobz, char* uplo, int* n, double* a, int* lda,
                    double* w, double* work, int* lwork, int* iwork,
                    int* liwork, int* info);
extern void dgesvd_(char* jobu, char* jobvt, int* m, int* n, double* a,
                    int* lda, double* s, double* u, int* ldu, double* vt,
                    int* ldvt, double* work, int* lwork, int* info);

/* --------------------------------------------------------------------------
 * Helper: query LAPACK workspace size
 * -------------------------------------------------------------------------- */
static int query_lwork_dgeqrf(int m, int n) {
    double work_query;
    int lwork = -1, info;
    dgeqrf_(&m, &n, NULL, &m, NULL, &work_query, &lwork, &info);
    return (int)work_query;
}

static int query_lwork_dorgqr(int m, int n, int k) {
    double work_query;
    int lwork = -1, info;
    dorgqr_(&m, &n, &k, NULL, &m, NULL, &work_query, &lwork, &info);
    return (int)work_query;
}

static int query_lwork_dsyevd(int n) {
    double work_query;
    int iwork_query;
    int lwork = -1, liwork = -1, info;
    char jobz = 'V', uplo = 'L';
    dsyevd_(&jobz, &uplo, &n, NULL, &n, NULL, &work_query, &lwork,
            &iwork_query, &liwork, &info);
    return (int)work_query;
}

static int query_liwork_dsyevd(int n) {
    double work_query;
    int iwork_query;
    int lwork = -1, liwork = -1, info;
    char jobz = 'V', uplo = 'L';
    dsyevd_(&jobz, &uplo, &n, NULL, &n, NULL, &work_query, &lwork,
            &iwork_query, &liwork, &info);
    return iwork_query;
}

static int query_lwork_dgesvd(int m, int n) {
    double work_query;
    int lwork = -1, info;
    char jobu = 'A', jobvt = 'A';
    dgesvd_(&jobu, &jobvt, &m, &n, NULL, &m, NULL, NULL, &m, NULL, &n,
            &work_query, &lwork, &info);
    return (int)work_query;
}

/* --------------------------------------------------------------------------
 * Cholesky: eigenhhlo_dpotrf
 *   in[0]  = A (n×n, double)
 *   out[0] = L or U (n×n, double, overwrites triangle)
 * -------------------------------------------------------------------------- */
void eigenhhlo_dpotrf(void* out, void** in,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int n = parse_int_kv(opaque, "n", 1);
    char uplo = parse_char_kv(opaque, "uplo", 'L');

    double* A = (double*)in[0];
    double* outA = (double*)out;  /* single result */

    /* Copy input to output (LAPACK overwrites in-place) */
    memcpy(outA, A, n * n * sizeof(double));

    int lda = n;
    int info;
    dpotrf_(&uplo, &n, outA, &lda, &info);
    (void)info; /* TODO: propagate error via status */
}

/* --------------------------------------------------------------------------
 * LU: eigenhhlo_dgetrf
 *   in[0]  = A (m×n, double)
 *   out[0] = LU (m×n, double)
 *   out[1] = pivots (min(m,n), int32)
 * -------------------------------------------------------------------------- */
void eigenhhlo_dgetrf(void* out, void** in,
                      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);

    double* A = (double*)in[0];
    void** outputs = (void**)out;
    double* outLU = (double*)outputs[0];
    int32_t* outPiv = (int32_t*)outputs[1];

    memcpy(outLU, A, m * n * sizeof(double));

    int lda = m;
    int info;
    int* piv = (int*)malloc((m < n ? m : n) * sizeof(int));
    dgetrf_(&m, &n, outLU, &lda, piv, &info);

    /* Copy pivots (1-based in LAPACK, convert to 0-based) */
    int min_mn = (m < n) ? m : n;
    for (int i = 0; i < min_mn; i++) {
        outPiv[i] = (int32_t)(piv[i] - 1);
    }
    free(piv);
    (void)info;
}

/* --------------------------------------------------------------------------
 * QR factorization: eigenhhlo_dgeqrf
 *   in[0]  = A (m×n, double)
 *   out[0] = A overwritten with R+reflectors (m×n, double)
 *   out[1] = tau (min(m,n), double)
 * -------------------------------------------------------------------------- */
void eigenhhlo_dgeqrf(void* out, void** in,
                      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);

    double* A = (double*)in[0];
    void** outputs = (void**)out;
    double* outA = (double*)outputs[0];
    double* outTau = (double*)outputs[1];

    memcpy(outA, A, m * n * sizeof(double));

    int lda = m;
    int min_mn = (m < n) ? m : n;
    int lwork = query_lwork_dgeqrf(m, n);
    double* work = (double*)malloc(lwork * sizeof(double));
    int info;

    dgeqrf_(&m, &n, outA, &lda, outTau, work, &lwork, &info);
    free(work);
    (void)info;
}

/* --------------------------------------------------------------------------
 * Generate Q from QR: eigenhhlo_dorgqr
 *   in[0]  = A (m×n, double, contains reflectors from dgeqrf)
 *   in[1]  = tau (k, double)
 *   out[0] = Q (m×m, double)  [full Q]
 * -------------------------------------------------------------------------- */
void eigenhhlo_dorgqr(void* out, void** in,
                      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);

    double* A = (double*)in[0];
    double* tau = (double*)in[1];
    double* outQ = (double*)out;  /* single result */

    /* Copy reflectors to output */
    memcpy(outQ, A, m * n * sizeof(double));

    int lda = m;
    int lwork = query_lwork_dorgqr(m, n, k);
    double* work = (double*)malloc(lwork * sizeof(double));
    int info;

    dorgqr_(&m, &n, &k, outQ, &lda, tau, work, &lwork, &info);
    free(work);
    (void)info;
}

/* --------------------------------------------------------------------------
 * Symmetric eigenvalue: eigenhhlo_dsyevd
 *   in[0]  = A (n×n, double, symmetric)
 *   out[0] = eigenvalues (n, double)
 *   out[1] = eigenvectors (n×n, double)
 * -------------------------------------------------------------------------- */
void eigenhhlo_dsyevd(void* out, void** in,
                      const char* opaque, size_t opaque_len,
                      XlaCustomCallStatus* status)
{
    int n = parse_int_kv(opaque, "n", 1);
    char uplo = parse_char_kv(opaque, "uplo", 'L');

    double* A = (double*)in[0];
    void** outputs = (void**)out;
    double* outW = (double*)outputs[0];
    double* outV = (double*)outputs[1];

    /* Copy A to eigenvectors (LAPACK overwrites in-place) */
    memcpy(outV, A, n * n * sizeof(double));

    int lda = n;
    int lwork = query_lwork_dsyevd(n);
    int liwork = query_liwork_dsyevd(n);
    double* work = (double*)malloc(lwork * sizeof(double));
    int* iwork = (int*)malloc(liwork * sizeof(int));
    int info;
    char jobz = 'V';

    dsyevd_(&jobz, &uplo, &n, outV, &lda, outW, work, &lwork,
            iwork, &liwork, &info);
    free(work);
    free(iwork);
    (void)info;
}

/* --------------------------------------------------------------------------
 * SVD: eigenhhlo_dgesvd
 *   in[0]  = A (m×n, double)
 *   out[0] = U (m×m or m×min, double)
 *   out[1] = S (min(m,n), double)
 *   out[2] = Vt (n×n or min×n, double)
 * -------------------------------------------------------------------------- */
void eigenhhlo_dgesvd(void* out, void** in,
                      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');

    double* A = (double*)in[0];
    void** outputs = (void**)out;
    double* outU  = (double*)outputs[0];
    double* outS  = (double*)outputs[1];
    double* outVt = (double*)outputs[2];

    /* Copy A to workspace (LAPACK overwrites) */
    double* awork = (double*)malloc(m * n * sizeof(double));
    memcpy(awork, A, m * n * sizeof(double));

    int lda = m;
    int ldu = m;
    int ldvt = n;
    int lwork = query_lwork_dgesvd(m, n);
    double* work = (double*)malloc(lwork * sizeof(double));
    int info;

    dgesvd_(&jobu, &jobvt, &m, &n, awork, &lda, outS, outU, &ldu,
            outVt, &ldvt, work, &lwork, &info);
    free(work);
    free(awork);
    (void)info;
}