// Blas.sgemm_w: C = op(A) op(W) with W a registry weight, on the CPU // BLAS. f: m, n, k, ta, tb, a (Array), id (U32). // The weight registry, defined once per program (the effects are spliced // into one C file, in some order): weights loaded from a file or copied // from an array live here, referenced by id, and a device copy is cached // beside them for the GPU path. #ifndef BEND_WEIGHTS #define BEND_WEIGHTS 1 #include #include #include #include typedef struct BendWeight { float* data; u32 n; void* dev; } BendWeight; static BendWeight* bend_weights = NULL; static u32 bend_weights_len = 0; static u32 bend_weights_cap = 0; static u32 bend_weight_add(float* data, u32 n) { if (bend_weights_len == bend_weights_cap) { bend_weights_cap = bend_weights_cap == 0 ? 16 : bend_weights_cap * 2; bend_weights = (BendWeight*)realloc(bend_weights, bend_weights_cap * sizeof(BendWeight)); } bend_weights[bend_weights_len].data = data; bend_weights[bend_weights_len].n = n; bend_weights[bend_weights_len].dev = NULL; bend_weights_len += 1; return bend_weights_len - 1; } static BendWeight* bend_weight_get(u32 id) { return id < bend_weights_len ? &bend_weights[id] : NULL; } // reads n float32 (little-endian, as numpy writes them) from path static float* bend_floats_read(const char* path, u32 n) { FILE* fp = fopen(path, "rb"); if (fp == NULL) { return NULL; } float* buf = (float*)malloc((size_t)n * sizeof(float)); size_t got = fread(buf, sizeof(float), n, fp); fclose(fp); if (got != n) { free(buf); return NULL; } return buf; } // a fresh Bend block of the smallest power-of-two size holding n floats, // the tail zero; answers 0 on allocation failure static Loc bend_block_floats(Env e, u32 n, Cls* cls_out) { Cls cls = 0; while ((1ull << cls) < (u64)n) { cls += 1; } if (cls > 31) { return 0; } Loc loc = heap_alloc(e, buf_wcls(cls)); if (err_seen(e.mem)) { return 0; } float* c = (float*)blk_ptr(e.mem, loc, 0); for (u64 i = n; i < (1ull << cls); i += 1) { c[i] = 0.0f; } *cls_out = cls; return loc; } // the platform CPU BLAS, opened at runtime (bend's build has fixed flags) typedef void (*bend_sgemm_fn)(int order, int ta, int tb, int m, int n, int k, float alpha, const float* a, int lda, const float* b, int ldb, float beta, float* c, int ldc); static bend_sgemm_fn bend_sgemm_get(void) { static bend_sgemm_fn fn = NULL; if (fn == NULL) { const char* names[] = { "/System/Library/Frameworks/Accelerate.framework/Accelerate", "libopenblas.so.0", "libopenblas.so", "libblas.so.3", "libblas.so", NULL }; for (int i = 0; names[i] != NULL && fn == NULL; i += 1) { void* h = dlopen(names[i], RTLD_NOW); if (h != NULL) { fn = (bend_sgemm_fn)dlsym(h, "cblas_sgemm"); } } } return fn; } // C (m x n) = op(A) op(B); A is m x k as stored (k x m when ta), B is // k x n as stored (n x k when tb); row-major throughout static void bend_sgemm_cpu(bend_sgemm_fn sgemm, u32 ta, u32 tb, u32 m, u32 n, u32 k, const float* a, const float* b, float* c) { sgemm(101, ta ? 112 : 111, tb ? 112 : 111, (int)m, (int)n, (int)k, 1.0f, a, ta ? (int)m : (int)k, b, tb ? (int)k : (int)n, 0.0f, c, (int)n); } #endif Term blas_sgemm_w_run(Env e, Term* f, IoWork* w) { u32 m = (u32)f[0]; u32 n = (u32)f[1]; u32 k = (u32)f[2]; u32 ta = (u32)f[3]; u32 tb = (u32)f[4]; BendWeight* wt = bend_weight_get((u32)f[6]); if (wt == NULL || wt->n < k * n) { return io_fail(e, 2, "Blas.sgemm_w: no such weight, or too small"); } bend_sgemm_fn sgemm = bend_sgemm_get(); if (sgemm == NULL) { return io_fail(e, 2, "Blas.sgemm_w: no BLAS found (Accelerate, OpenBLAS)"); } const float* a = (const float*)blk_ptr(e.mem, term_loc(f[5]), 0); Cls cls = 0; Loc loc = bend_block_floats(e, m * n, &cls); if (loc == 0) { return io_fail(e, 2, "Blas.sgemm_w: out of memory"); } if (m > 0 && n > 0 && k > 0) { bend_sgemm_cpu(sgemm, ta, tb, m, n, k, a, wt->data, (float*)blk_ptr(e.mem, loc, 0)); } return io_done(e, term_blk(false, cls, loc)); } static void __attribute__((constructor)) blas_sgemm_w_use(void) { io_eff(CID_BLAS_SGEMM_W, blas_sgemm_w_run, 0); }