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