+ New

utensil

Public
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