← Commits · 4350af2a
4350af2a8cdbec426849ab5200632b3fb2b3c617
diff --git a/Makefile b/Makefile
index ec84f5a..52e7f5c 100644
--- a/Makefile
+++ b/Makefile
@@ -3,7 +3,8 @@ LDFLAGS ?= -lm
METAL := $(shell test -d /System/Library/Frameworks/Metal.framework && echo 1 || echo 0)
ifeq ($(METAL),1)
- LDFLAGS += -framework Metal -framework MetalPerformanceShaders -framework Foundation -framework Accelerate
+ CFLAGS += -DACCELERATE_NEW_LAPACK
+ LDFLAGS += -framework Metal -framework MetalPerformanceShaders -framework Foundation -framework Accelerate
endif
all:
diff --git a/test.c b/test.c
index 961d613..c03e05c 100644
--- a/test.c
+++ b/test.c
@@ -47,9 +47,9 @@ void test_reshape(void) {
ut_free(a);
}
-void test_elementwise(void) {
- ut_tensor* a = ut_from_data(1, (int[]){5}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f}, UT_METAL);
- ut_tensor* b = ut_from_data(1, (int[]){5}, (float[]){5.f, 4.f, 3.f, 2.f, 1.f}, UT_METAL);
+void test_elementwise(ut_dev dev) {
+ ut_tensor* a = ut_from_data(1, (int[]){5}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f}, dev);
+ ut_tensor* b = ut_from_data(1, (int[]){5}, (float[]){5.f, 4.f, 3.f, 2.f, 1.f}, dev);
// unary
{
ut_tensor* c = ut_neg(a);
@@ -118,10 +118,46 @@ void test_elementwise(void) {
ut_free(b);
}
+static void test_matmul_2d(ut_dev dev) {
+ // |1 2| | 7 8 9| | 58 64|
+ // |3 4| x |10 11 12| = |139 154|
+ // |5 6|
+ //
+ ut_tensor* a = ut_from_data(2, (int[]){2, 3}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f}, dev);
+ ut_tensor* b = ut_from_data(2, (int[]){3, 2}, (float[]){7.f, 8.f, 9.f, 10.f, 11.f, 12.f}, dev);
+ ut_tensor* c = ut_matmul(a, b);
+ ut_sync_cpu(c);
+ assert(c->shape.ndim == 2 && c->shape.shape[0] == 2 && c->shape.shape[1] == 2);
+ assert_data(c, ((float[]){58.f, 64.f, 139.f, 154.f}), 1e-6f);
+ ut_free(a);
+ ut_free(b);
+ ut_free(c);
+}
+
+static void test_matmul_3d(ut_dev dev) {
+ ut_tensor* a =
+ ut_from_data(3, (int[]){2, 2, 3},
+ (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 10.f, 11.f, 12.f}, dev);
+ ut_tensor* b = ut_from_data(
+ 3, (int[]){2, 3, 2},
+ (float[]){13.f, 14.f, 15.f, 16.f, 17.f, 18.f, 19.f, 20.f, 21.f, 22.f, 23.f, 24.f}, dev);
+ ut_tensor* c = ut_matmul(a, b);
+ ut_sync_cpu(c);
+ assert_data(c, ((float[]){94.f, 100.f, 229.f, 244.f, 508.f, 532.f, 697.f, 730.f}), 1e-6f);
+ ut_free(a);
+ ut_free(b);
+ ut_free(c);
+}
+
int main() {
test_shape();
test_lifetime();
test_reshape();
- test_elementwise();
+ test_elementwise(UT_CPU);
+ test_elementwise(UT_METAL);
+ test_matmul_2d(UT_CPU);
+ test_matmul_2d(UT_METAL);
+ test_matmul_3d(UT_CPU);
+ test_matmul_3d(UT_METAL);
return 0;
}
diff --git a/utensil.h b/utensil.h
index b500d48..d463e0e 100644
--- a/utensil.h
+++ b/utensil.h
@@ -69,7 +69,7 @@ static inline const char* _c0(void* o, const char* s) {
static const char* _mtl_src =
"#include <metal_stdlib>\nusing namespace metal;\n"
- /* unary */
+ // unary
"kernel void uneg(device const float*i,device float*o,constant int&n,"
"uint idx[[thread_position_in_grid]]){o[idx]=-i[idx];}\n"
"kernel void urelu(device const float*i,device float*o,constant int&n,"
@@ -81,7 +81,7 @@ static const char* _mtl_src =
"kernel void uexp(device const float*i,device float*o,constant int&n,"
"uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=exp(i[idx]);}\n"
- /* binary */
+ // binary
"kernel void badd(device const float*a,device const float*b,device float*o,"
"constant int&n,uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=a[idx]+b[idx];}\n"
"kernel void bsub(device const float*a,device const float*b,device float*o,"
@@ -91,11 +91,21 @@ static const char* _mtl_src =
"kernel void bdiv(device const float*a,device const float*b,device float*o,"
"constant int&n,uint idx[[thread_position_in_grid]]){if((int)idx<n)o[idx]=a[idx]/b[idx];}\n"
- /* scalar multiply: args = {float s, int n} packed as 8 bytes */
+ // scalar multiply: args = {float s, int n} packed as 8 bytes
"struct ScaleArgs{ float s; int n; };\n"
"kernel void bscale(device const float*a,device float*o,"
"constant ScaleArgs&args,uint idx[[thread_position_in_grid]])"
"{if((int)idx<args.n)o[idx]=a[idx]*args.s;}\n"
+ // batched matmul: c[b,m,n] = sum_k a[b,m,k]*b[b,k,n]
+ // p = {B,M,N,K}, dispatch B*M*N threads (flat 1D)
+ "kernel void bmatmul(device const float*a,device const float*b_,device float*c,"
+ "constant int*p,uint idx[[thread_position_in_grid]]){"
+ "int M=p[1],N=p[2],K=p[3];"
+ "int tot=p[0]*M*N;if((int)idx>=tot)return;"
+ "int n=(int)idx%N,t=(int)idx/N,m=t%M,bat=t/M;"
+ "float s=0;int ao=bat*M*K+m*K,bo=bat*K*N+n;"
+ "for(int k=0;k<K;k++)s+=a[ao+k]*b_[bo+k*N];"
+ "c[idx]=s;}\n"
"";
#define _MTL_MAX_PL 32 // maximum number of pipeline states
@@ -206,6 +216,47 @@ static void _mtl_dispatch(_mtl_ctx_t* c, const char* kern, void** bufs, int nbuf
_v0(enc, "endEncoding");
}
+static void _mtl_matmul(_mtl_ctx_t* ctx, void* a, void* b, void* res, int m, int n, int k, bool ta,
+ bool tb) {
+ int ra = ta ? k : m, ca = ta ? m : k;
+ int rb = tb ? n : k, cb = tb ? k : n;
+ long s = (long)sizeof(float);
+ unsigned long mpsf32 = 0x10000020UL; // MPSDataTypeFloat32
+ void* dA =
+ _m4l(objc_getClass("MPSMatrixDescriptor"),
+ "matrixDescriptorWithRows:columns:rowBytes:dataType:", ra, ca, ca * s, (long)mpsf32);
+ void* dB =
+ _m4l(objc_getClass("MPSMatrixDescriptor"),
+ "matrixDescriptorWithRows:columns:rowBytes:dataType:", rb, cb, cb * s, (long)mpsf32);
+ void* dC = _m4l(objc_getClass("MPSMatrixDescriptor"),
+ "matrixDescriptorWithRows:columns:rowBytes:dataType:", m, n, n * s, (long)mpsf32);
+ void* mA = ((id (*)(id, SEL, id, id))objc_msgSend)(
+ ((id (*)(id, SEL))objc_msgSend)((id)objc_getClass("MPSMatrix"), sel_getUid("alloc")),
+ sel_getUid("initWithBuffer:descriptor:"), (id)a, (id)dA);
+ void* mB = ((id (*)(id, SEL, id, id))objc_msgSend)(
+ ((id (*)(id, SEL))objc_msgSend)((id)objc_getClass("MPSMatrix"), sel_getUid("alloc")),
+ sel_getUid("initWithBuffer:descriptor:"), (id)b, (id)dB);
+ void* mC = ((id (*)(id, SEL, id, id))objc_msgSend)(
+ ((id (*)(id, SEL))objc_msgSend)((id)objc_getClass("MPSMatrix"), sel_getUid("alloc")),
+ sel_getUid("initWithBuffer:descriptor:"), (id)res, (id)dC);
+ void* mm = ((id (*)(id, SEL, id, bool, bool, unsigned long, unsigned long, unsigned long, double,
+ double))objc_msgSend)(
+ ((id (*)(id, SEL))objc_msgSend)((id)objc_getClass("MPSMatrixMultiplication"),
+ sel_getUid("alloc")),
+ sel_getUid("initWithDevice:transposeLeft:transposeRight:"
+ "resultRows:resultColumns:interiorColumns:alpha:beta:"),
+ (id)ctx->device, ta, tb, (unsigned long)m, (unsigned long)n, (unsigned long)k, 1.0, 0.0);
+ void* cmd = _mtl_get_cmd(ctx);
+ ((void (*)(id, SEL, id, id, id, id))objc_msgSend)(
+ (id)mm, sel_getUid("encodeToCommandBuffer:leftMatrix:rightMatrix:resultMatrix:"), (id)cmd,
+ (id)mA, (id)mB, (id)mC);
+ // no commit — work is batched into pending_cmd
+ if (mm) _v0(mm, "release");
+ if (mA) _v0(mA, "release");
+ if (mB) _v0(mB, "release");
+ if (mC) _v0(mC, "release");
+}
+
static void* _mtl_buf_alloc(_mtl_ctx_t* c, size_t bytes) {
return _m2ll(c->device, "newBufferWithLength:options:", (long)bytes, 0L);
}
@@ -466,4 +517,74 @@ ut_tensor* ut_scale(ut_tensor* a, float s) {
return out;
}
+// ========================================================
+// Matmul
+// ========================================================
+static void gemm(const float* a, const float* b, float* c, int m, int n, int k, bool ta, bool tb) {
+#ifdef __APPLE__
+ cblas_sgemm(CblasRowMajor, ta ? CblasTrans : CblasNoTrans, tb ? CblasTrans : CblasNoTrans, m, n,
+ k, 1.0f, a, ta ? m : k, b, tb ? k : n, 0.0f, c, n);
+#else
+ for (int i = 0; i < m; i++)
+ for (int j = 0; j < n; j++) {
+ float sum = 0;
+ for (int l = 0; l < k; l++)
+ sum += (ta ? a[l * m + i] : a[i * k + l]) * (tb ? b[j * k + l] : b[l * n + j]);
+ c[i * n + j] = sum;
+ }
+#endif
+}
+
+ut_tensor* ut_matmul(ut_tensor* a, ut_tensor* b) {
+ // 2Dx2D: MPS path
+ if (a->shape.ndim == 2 && b->shape.ndim == 2) {
+ int m = a->shape.shape[0], k = a->shape.shape[1], n = b->shape.shape[1];
+ if (b->shape.shape[0] != k) return NULL;
+ ut_dev dev = (a->dev == UT_METAL || b->dev == UT_METAL) ? UT_METAL : UT_CPU;
+ ut_to_device(a, dev);
+ ut_to_device(b, dev);
+ int cd[2] = {m, n};
+ ut_tensor* c = ut_alloc(2, cd, dev);
+ _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+ if (dev == UT_METAL && mc) {
+ ut_sync_gpu(a);
+ ut_sync_gpu(b);
+ _mtl_matmul(mc, a->gpu_buf, b->gpu_buf, c->gpu_buf, m, n, k, false, false);
+ c->dirty_cpu = true;
+ } else {
+ ut_sync_cpu(a);
+ ut_sync_cpu(b);
+ gemm(a->data, b->data, c->data, m, n, k, false, false);
+ }
+ return c;
+ }
+ // batched [B,M,K] x [B,K,N]: GPU via bmatmul kernel, CPU fallback
+ if (a->shape.ndim == 3 && b->shape.ndim == 3) {
+ int B = a->shape.shape[0], m = a->shape.shape[1], k = a->shape.shape[2];
+ if (b->shape.shape[0] != B || b->shape.shape[1] != k) return NULL;
+ int n = b->shape.shape[2];
+ int cd[3] = {B, m, n};
+ ut_dev dev = (a->dev == UT_METAL || b->dev == UT_METAL) ? UT_METAL : UT_CPU;
+ _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+ if (dev == UT_METAL && mc) {
+ ut_to_device(a, dev);
+ ut_to_device(b, dev);
+ ut_tensor* c = ut_alloc(3, cd, UT_METAL);
+ int p[4] = {B, m, n, k};
+ void* bufs[3] = {a->gpu_buf, b->gpu_buf, c->gpu_buf};
+ // dispatch B*M*N threads via a 3D grid encoded as 1D
+ _mtl_dispatch(mc, "bmatmul", bufs, 3, p, (int)sizeof(p), B * m * n);
+ c->dirty_cpu = true;
+ return c;
+ }
+ ut_sync_cpu(a);
+ ut_sync_cpu(b);
+ ut_tensor* c = ut_alloc(3, cd, UT_CPU);
+ for (int bi = 0; bi < B; bi++)
+ gemm(a->data + bi * m * k, b->data + bi * k * n, c->data + bi * m * n, m, n, k, false, false);
+ return c;
+ }
+ return NULL;
+}
+
#endif // UTENSIL_H