+ New

utensil

Public
dd4b9aebc3362ddc64fa6a3129c87c5f55ac3860
diff --git a/test.c b/test.c
index 42a35a7..33ecafa 100644
--- a/test.c
+++ b/test.c
@@ -137,6 +137,14 @@ static void test_edge_activations(ut_dev dev) {
     assert_data(c, ((float[]){0.f, 0.f, 0.f, 3.f, 4.f, 8.f}), 1e-6f);
     ut_free(c);
   }
+  {
+    // gradient passes through only where 0 < a < 6: a=3,4 -> 1; everything else -> 0
+    ut_tensor* go = ut_from_data(1, (int[]){6}, (float[]){1.f, 1.f, 1.f, 1.f, 1.f, 1.f}, dev);
+    ut_tensor* gi = ut_relu6_backward(go, a);
+    ut_sync_cpu(gi);
+    assert_data(gi, ((float[]){0.f, 0.f, 0.f, 1.f, 1.f, 0.f}), 1e-6f);
+    ut_free_all(go, gi);
+  }
   ut_free(a);
 }
 
@@ -505,6 +513,70 @@ static void test_conv2d_backward(ut_dev dev) {
   ut_conv2d_free(&l);
 }
 
+static void test_dwconv2d_forward(ut_dev dev) {
+  // N=1,C=2,kh=kw=2,s=1,p=0. ch0=[[1,2,3],[4,5,6],[7,8,9]] W=[1,0;0,-1] (same as
+  // the plain conv2d fixture, restricted to one channel: out=[-4,-4,-4,-4]).
+  // ch1=[[9,8,7],[6,5,4],[3,2,1]] W=[0,1;1,0]: out=[14,12,8,6] (uses its OWN
+  // input, not ch0's -- the whole point of depthwise being channel-independent).
+  ut_dwconv2d l = ut_dwconv2d_alloc(2, 2, 2, 1, 0, true, dev);
+  ut_free(l.weight);
+  l.weight = ut_from_data(3, (int[]){2, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f, 0.f, 1.f, 1.f, 0.f},
+                          dev);
+  ut_free(l.bias);
+  l.bias = ut_from_data(1, (int[]){2}, (float[]){100.f, 1000.f}, dev);
+
+  ut_tensor* x = ut_from_data(
+      4, (int[]){1, 2, 3, 3},
+      (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 9.f, 8.f, 7.f, 6.f, 5.f, 4.f, 3.f, 2.f,
+                1.f},
+      dev);
+  ut_tensor* out = ut_dwconv2d_forward(&l, x, NULL);
+  ut_sync_cpu(out);
+  assert(out->shape.ndim == 4 && out->shape.shape[1] == 2 && out->shape.shape[2] == 2 &&
+         out->shape.shape[3] == 2);
+  assert_data(out, ((float[]){96.f, 96.f, 96.f, 96.f, 1014.f, 1012.f, 1008.f, 1006.f}), 1e-4f);
+
+  ut_free_all(out, x);
+  ut_dwconv2d_free(&l);
+}
+
+static void test_dwconv2d_backward(ut_dev dev) {
+  // same x/weight as forward, no bias; grad_out one-hot on channel0's first
+  // output position -> channel1 (and thus its dW/dx) must stay exactly zero.
+  ut_dwconv2d l = ut_dwconv2d_alloc(2, 2, 2, 1, 0, false, dev);
+  ut_free(l.weight);
+  l.weight = ut_from_data(3, (int[]){2, 2, 2}, (float[]){1.f, 0.f, 0.f, -1.f, 0.f, 1.f, 1.f, 0.f},
+                          dev);
+
+  ut_tensor* x = ut_from_data(
+      4, (int[]){1, 2, 3, 3},
+      (float[]){1.f, 2.f, 3.f, 4.f, 5.f, 6.f, 7.f, 8.f, 9.f, 9.f, 8.f, 7.f, 6.f, 5.f, 4.f, 3.f, 2.f,
+                1.f},
+      dev);
+  ut_dwconv2d_cache c;
+  ut_tensor* out = ut_dwconv2d_forward(&l, x, &c);
+  ut_free(out);
+
+  ut_tensor* go = ut_from_data(4, (int[]){1, 2, 2, 2},
+                               (float[]){1.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f}, dev);
+  ut_tensor* dW = ut_alloc(1, (int[]){2 * 2 * 2}, dev);
+  memset(dW->data, 0, 8 * sizeof(float));
+
+  ut_tensor* dx = ut_dwconv2d_backward(&l, &c, go, dW, NULL);
+  ut_sync_cpu(dW);
+  ut_sync_cpu(dx);
+  // channel0: identical to the plain conv2d backward test's dW/dx
+  assert_data(dW, ((float[]){1.f, 2.f, 4.f, 5.f, 0.f, 0.f, 0.f, 0.f}), 1e-4f);
+  assert_data(dx,
+              ((float[]){1.f, 0.f, 0.f, 0.f, -1.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f,
+                         0.f, 0.f, 0.f}),
+              1e-4f);
+
+  ut_free_all(dx, dW, go, x);
+  ut_dwconv2d_cache_free(&c);
+  ut_dwconv2d_free(&l);
+}
+
 static void test_mse(ut_dev dev) {
   // pred=[[1,2],[3,4]], target=[[1,0],[3,6]]
   // diffs=[0,2,0,-2]    mse = mean(diffs^2) = (0+4+0+4)/4 = 2.0
@@ -624,6 +696,11 @@ int main() {
   test_conv2d_backward(UT_CPU);
   test_conv2d_backward(UT_METAL);
 
+  test_dwconv2d_forward(UT_CPU);
+  test_dwconv2d_forward(UT_METAL);
+  test_dwconv2d_backward(UT_CPU);
+  test_dwconv2d_backward(UT_METAL);
+
   test_sgd_momentum();
   test_adam();
   return 0;
diff --git a/utensil.h b/utensil.h
index be3cd8e..2eae1e5 100644
--- a/utensil.h
+++ b/utensil.h
@@ -69,6 +69,16 @@ typedef struct ut_conv2d_cache {
   ut_tensor* col;
 } ut_conv2d_cache;
 
+typedef struct ut_dwconv2d {
+  ut_tensor* weight;  // [C, kh, kw] — one filter per channel, no cross-channel mixing
+  ut_tensor* bias;    // [C]
+  int c, kh, kw, stride, pad;
+} ut_dwconv2d;
+
+typedef struct ut_dwconv2d_cache {
+  ut_tensor* input;
+} ut_dwconv2d_cache;
+
 typedef struct ut_batchnorm2d {
   ut_tensor* weight;        // [C] gain
   ut_tensor* bias;          // [C] shift
@@ -203,6 +213,10 @@ static const char* _mtl_src =
     "kernel void relu_bwd(device const float*go,device const float*fwd,"
     "device float*gi,constant int&n,"
     "uint idx[[thread_position_in_grid]]){if((int)idx<n)gi[idx]=fwd[idx]>0.f?go[idx]:0.f;}"
+    "kernel void relu6_bwd(device const float*go,device const float*fwd,"
+    "device float*gi,constant int&n,"
+    "uint idx[[thread_position_in_grid]]){if((int)idx<n)"
+    "gi[idx]=(fwd[idx]>0.f&&fwd[idx]<6.f)?go[idx]:0.f;}"
     // transpose dims 0<->1: p={A,B,inner}
     "kernel void transpose_01("
     "device const float*in,device float*out,constant int*p,"
@@ -440,6 +454,71 @@ static const char* _mtl_src =
     "int ih=oh*s-pad+khh,iw=ow*s-pad+kww;"
     "float v=0;if(ih>=0&&ih<H&&iw>=0&&iw<W)v=x[((n*C+c)*H+ih)*W+iw];"
     "col[(c*kh*kw+khh*kw+kww)*N*Ho*Wo+n*Ho*Wo+oh*Wo+ow]=v;}\n"
+    // depthwise conv2d: each channel convolved independently with its own
+    // kh x kw filter, no cross-channel mixing -- im2col+matmul can't express
+    // this (it always contracts over all channels), so it's a direct kernel.
+    // p = {N,C,H,W,kh,kw,stride,pad,Ho,Wo}, weight is [C,kh,kw]
+    "kernel void dwconv2d_fwd(device const float*x,device const float*w,device float*out,"
+    "constant int*p,uint idx[[thread_position_in_grid]]){"
+    "int C=p[1],H=p[2],W=p[3],kh=p[4],kw=p[5],s=p[6],pad=p[7],Ho=p[8],Wo=p[9];"
+    "int tot=p[0]*C*Ho*Wo;if((int)idx>=tot)return;"
+    "int t=(int)idx,ow=t%Wo;t/=Wo;int oh=t%Ho;t/=Ho;int c=t%C;t/=C;int n=t;"
+    "float sum=0;"
+    "for(int khh=0;khh<kh;khh++)for(int kww=0;kww<kw;kww++){"
+    "int ih=oh*s-pad+khh,iw=ow*s-pad+kww;"
+    "if(ih<0||ih>=H||iw<0||iw>=W)continue;"
+    "sum+=x[((n*C+c)*H+ih)*W+iw]*w[(c*kh+khh)*kw+kww];}"
+    "out[idx]=sum;}\n"
+    // depthwise conv2d backward, dx (gather, no atomics — mirrors col2im_k)
+    "kernel void dwconv2d_bwd_dx(device const float*go,device const float*w,device float*dx,"
+    "constant int*p,uint idx[[thread_position_in_grid]]){"
+    "int C=p[1],H=p[2],W=p[3],kh=p[4],kw=p[5],s=p[6],pad=p[7],Ho=p[8],Wo=p[9];"
+    "int tot=p[0]*C*H*W;if((int)idx>=tot)return;"
+    "int t=(int)idx,iw=t%W;t/=W;int ih=t%H;t/=H;int c=t%C;t/=C;int n=t;"
+    "float sum=0;"
+    "for(int khh=0;khh<kh;khh++)for(int kww=0;kww<kw;kww++){"
+    "int oh_n=ih+pad-khh,ow_n=iw+pad-kww;"
+    "if(oh_n%s!=0||ow_n%s!=0)continue;"
+    "int oh=oh_n/s,ow=ow_n/s;"
+    "if(oh<0||oh>=Ho||ow<0||ow>=Wo)continue;"
+    "sum+=go[((n*C+c)*Ho+oh)*Wo+ow]*w[(c*kh+khh)*kw+kww];}"
+    "dx[idx]=sum;}\n"
+    // depthwise conv2d backward, dW: one threadgroup per (c,khh,kww), reducing
+    // over N*Ho*Wo. Writes dW_partial[C*kh*kw] (caller adds it into the real
+    // accumulator) so only that tiny buffer needs to come back to CPU.
+    "kernel void dwconv2d_bwd_dw(device const float*go,device const float*x,device float*dwp,"
+    "constant int*p,"
+    "uint gid[[threadgroup_position_in_grid]],"
+    "uint lid[[thread_position_in_threadgroup]],"
+    "uint tgs[[threads_per_threadgroup]]){"
+    "int C=p[1],H=p[2],W=p[3],kh=p[4],kw=p[5],s=p[6],pad=p[7],Ho=p[8],Wo=p[9];"
+    "int kww=(int)gid%kw,t0=(int)gid/kw,khh=t0%kh,c=t0/kh;"
+    "int M=p[0]*Ho*Wo;"
+    "threadgroup float sh[1024];"
+    "float ls=0;"
+    "for(int i=(int)lid;i<M;i+=(int)tgs){"
+    "int t=i,ow=t%Wo;t/=Wo;int oh=t%Ho;t/=Ho;int n=t;"
+    "int ih=oh*s-pad+khh,iw=ow*s-pad+kww;"
+    "if(ih<0||ih>=H||iw<0||iw>=W)continue;"
+    "ls+=go[((n*C+c)*Ho+oh)*Wo+ow]*x[((n*C+c)*H+ih)*W+iw];}"
+    "sh[lid]=ls;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint ss=tgs/2;ss>0;ss>>=1){if(lid<ss)sh[lid]+=sh[lid+ss];"
+    "threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "if(lid==0)dwp[gid]=sh[0];}\n"
+    // depthwise conv2d backward, db: one threadgroup per channel, reducing go
+    // over N*Ho*Wo (same reduction shape as bn2d_bwd's sum_go)
+    "kernel void dwconv2d_bwd_db(device const float*go,device float*dbp,"
+    "constant int*p,"
+    "uint gid[[threadgroup_position_in_grid]],"
+    "uint lid[[thread_position_in_threadgroup]],"
+    "uint tgs[[threads_per_threadgroup]]){"
+    "int C=p[1],HW=p[8]*p[9],c=(int)gid,M=p[0]*HW;"
+    "threadgroup float sh[1024];"
+    "float ls=0;for(int i=(int)lid;i<M;i+=(int)tgs){int n=i/HW,hw=i%HW;ls+=go[(n*C+c)*HW+hw];}"
+    "sh[lid]=ls;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint ss=tgs/2;ss>0;ss>>=1){if(lid<ss)sh[lid]+=sh[lid+ss];"
+    "threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "if(lid==0)dbp[c]=sh[0];}\n"
     "";
 
 #define _MTL_MAX_PL 32  // maximum number of pipeline states
@@ -1129,6 +1208,25 @@ ut_tensor* ut_relu_backward(ut_tensor* grad_out, ut_tensor* fwd_input) {
   }
   return gi;
 }
+
+ut_tensor* ut_relu6_backward(ut_tensor* grad_out, ut_tensor* fwd_input) {
+  _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+  ut_tensor* gi = ut_alloc(grad_out->shape.ndim, grad_out->shape.shape, grad_out->dev);
+  if (grad_out->dev == UT_METAL && mc) {
+    ut_sync_gpu(grad_out);
+    ut_to_device(fwd_input, UT_METAL);
+    int n = grad_out->shape.nelem;
+    _mtl_dispatch(mc, "relu6_bwd", (void*[]){grad_out->gpu_buf, fwd_input->gpu_buf, gi->gpu_buf}, 3,
+                  &n, sizeof(int), n);
+    gi->dirty_cpu = true;
+  } else {
+    ut_sync_cpu(grad_out);
+    ut_sync_cpu(fwd_input);
+    for (int i = 0; i < grad_out->shape.nelem; i++)
+      gi->data[i] = (fwd_input->data[i] > 0 && fwd_input->data[i] < 6) ? grad_out->data[i] : 0.f;
+  }
+  return gi;
+}
 // =========================================================
 // LayerNorm
 // =========================================================
@@ -1543,8 +1641,36 @@ ut_tensor* ut_conv2d_backward(ut_conv2d* l, ut_conv2d_cache* cache, ut_tensor* g
   int H = cache->input->shape.shape[2], W = cache->input->shape.shape[3];
   int Ho = (H + 2 * l->pad - l->kh) / l->stride + 1;
   int Wo = (W + 2 * l->pad - l->kw) / l->stride + 1;
-  ut_sync_cpu(grad_out);
-  ut_sync_cpu(dW);
+
+  if (l->bias && db) {
+    _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+    if (mc && grad_out->dev == UT_METAL) {
+      // reuse dwconv2d_bwd_db's per-channel reduction over [N,C,Ho,Wo] -- only
+      // reads p[0]=N, p[1]=C, p[8]=Ho, p[9]=Wo, so it applies unchanged here
+      ut_sync_gpu(grad_out);
+      int params[10] = {N, l->out_c, H, W, l->kh, l->kw, l->stride, l->pad, Ho, Wo};
+      int M = N * Ho * Wo;
+      int tgsize = M < 64 ? 32 : (M < 256 ? 64 : (M < 512 ? 128 : 256));
+      if (tgsize > 1024) tgsize = 1024;
+      ut_tensor* dbp = ut_alloc(1, (int[]){l->out_c}, UT_METAL);
+      _mtl_dispatch_rows(mc, "dwconv2d_bwd_db", (void*[]){grad_out->gpu_buf, dbp->gpu_buf}, 2,
+                         params, (int)sizeof(params), l->out_c, tgsize);
+      dbp->dirty_cpu = true;
+      ut_sync_cpu(dbp);
+      ut_sync_cpu(db);
+      for (int c = 0; c < l->out_c; c++) db->data[c] += dbp->data[c];
+      db->dirty_gpu = true;
+      ut_free(dbp);
+    } else {
+      ut_sync_cpu(grad_out);
+      ut_sync_cpu(db);
+      for (int n = 0; n < N; n++)
+        for (int oc = 0; oc < l->out_c; oc++)
+          for (int h = 0; h < Ho; h++)
+            for (int w = 0; w < Wo; w++)
+              db->data[oc] += grad_out->data[((n * l->out_c + oc) * Ho + h) * Wo + w];
+    }
+  }
 
   // go_2d: [out_c, N*Ho*Wo] — transpose + reshape grad_out [N,out_c,Ho,Wo]
   ut_tensor* go_t = ut_transpose(grad_out, 0, 1);       // [out_c, N, Ho, Wo]
@@ -1556,16 +1682,6 @@ ut_tensor* ut_conv2d_backward(ut_conv2d* l, ut_conv2d_cache* cache, ut_tensor* g
   for (int i = 0; i < dW->shape.nelem; i++) dW->data[i] += dWb->data[i];
   ut_free(dWb);
 
-  if (l->bias && db) {
-    ut_sync_cpu(db);
-    ut_sync_cpu(grad_out);
-    for (int n = 0; n < N; n++)
-      for (int oc = 0; oc < l->out_c; oc++)
-        for (int h = 0; h < Ho; h++)
-          for (int w = 0; w < Wo; w++)
-            db->data[oc] += grad_out->data[((n * l->out_c + oc) * Ho + h) * Wo + w];
-  }
-
   ut_tensor w_view = *l->weight;
   w_view.shape = ut_shape_new(2, (int[]){l->out_c, l->in_c * l->kh * l->kw});
   w_view.rc = 0x7fffffff;
@@ -1587,6 +1703,166 @@ void ut_conv2d_free(ut_conv2d* l) {
   l->weight = l->bias = NULL;
 }
 
+// =========================================================
+// Depthwise Conv2D
+// =========================================================
+ut_dwconv2d ut_dwconv2d_alloc(int c, int kh, int kw, int stride, int pad, bool bias, ut_dev dev) {
+  ut_dwconv2d l = {.c = c, .kh = kh, .kw = kw, .stride = stride, .pad = pad};
+  float std = sqrtf(2.f / (float)(kh * kw));
+  l.weight = ut_randn(3, (int[]){c, kh, kw}, 0.f, std, dev);
+  if (bias) l.bias = ut_alloc(1, (int[]){c}, dev);
+  return l;
+}
+
+// x: [N,C,H,W] -> [N,C,Ho,Wo], each channel convolved with its own kh x kw
+// filter (no cross-channel mixing) -- the depthwise half of a depthwise-
+// separable conv; follow with a 1x1 ut_conv2d for the pointwise half.
+ut_tensor* ut_dwconv2d_forward(ut_dwconv2d* l, ut_tensor* x, ut_dwconv2d_cache* cache) {
+  int N = x->shape.shape[0], C = l->c, H = x->shape.shape[2], W = x->shape.shape[3];
+  int Ho = (H + 2 * l->pad - l->kh) / l->stride + 1;
+  int Wo = (W + 2 * l->pad - l->kw) / l->stride + 1;
+  int params[10] = {N, C, H, W, l->kh, l->kw, l->stride, l->pad, Ho, Wo};
+
+  _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+  ut_tensor* out;
+  if (mc && x->dev == UT_METAL) {
+    ut_sync_gpu(x);
+    ut_to_device(l->weight, UT_METAL);
+    out = ut_alloc(4, (int[]){N, C, Ho, Wo}, UT_METAL);
+    _mtl_dispatch(mc, "dwconv2d_fwd", (void*[]){x->gpu_buf, l->weight->gpu_buf, out->gpu_buf}, 3,
+                  params, (int)sizeof(params), N * C * Ho * Wo);
+    out->dirty_cpu = true;
+  } else {
+    ut_sync_cpu(x);
+    ut_sync_cpu(l->weight);
+    out = ut_alloc(4, (int[]){N, C, Ho, Wo}, UT_CPU);
+    for (int n = 0; n < N; n++)
+      for (int c = 0; c < C; c++)
+        for (int oh = 0; oh < Ho; oh++)
+          for (int ow = 0; ow < Wo; ow++) {
+            float sum = 0;
+            for (int khh = 0; khh < l->kh; khh++)
+              for (int kww = 0; kww < l->kw; kww++) {
+                int ih = oh * l->stride - l->pad + khh, iw = ow * l->stride - l->pad + kww;
+                if (ih < 0 || ih >= H || iw < 0 || iw >= W) continue;
+                sum += x->data[((n * C + c) * H + ih) * W + iw] *
+                       l->weight->data[(c * l->kh + khh) * l->kw + kww];
+              }
+            out->data[((n * C + c) * Ho + oh) * Wo + ow] = sum;
+          }
+    out->dirty_gpu = true;
+  }
+
+  if (l->bias) {
+    if (out->dev == UT_METAL && mc) {
+      ut_sync_gpu(out);
+      ut_to_device(l->bias, UT_METAL);
+      int bparams[3] = {N, C, Ho * Wo};
+      _mtl_dispatch(mc, "bias_add", (void*[]){out->gpu_buf, l->bias->gpu_buf}, 2, bparams,
+                    (int)sizeof(bparams), N * C * Ho * Wo);
+      out->dirty_cpu = true;
+    } else {
+      ut_sync_cpu(out);
+      ut_sync_cpu(l->bias);
+      for (int n = 0; n < N; n++)
+        for (int c = 0; c < C; c++)
+          for (int h = 0; h < Ho; h++)
+            for (int w = 0; w < Wo; w++)
+              out->data[((n * C + c) * Ho + h) * Wo + w] += l->bias->data[c];
+    }
+  }
+
+  if (cache) cache->input = ut_retain(x);
+  return out;
+}
+
+ut_tensor* ut_dwconv2d_backward(ut_dwconv2d* l, ut_dwconv2d_cache* cache, ut_tensor* grad_out,
+                                ut_tensor* dW, ut_tensor* db) {
+  ut_tensor* x = cache->input;
+  int N = x->shape.shape[0], C = l->c, H = x->shape.shape[2], W = x->shape.shape[3];
+  int Ho = (H + 2 * l->pad - l->kh) / l->stride + 1;
+  int Wo = (W + 2 * l->pad - l->kw) / l->stride + 1;
+  int params[10] = {N, C, H, W, l->kh, l->kw, l->stride, l->pad, Ho, Wo};
+
+  _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+  if (mc && grad_out->dev == UT_METAL) {
+    ut_sync_gpu(grad_out);
+    ut_to_device(x, UT_METAL);
+    ut_to_device(l->weight, UT_METAL);
+
+    int M = N * Ho * Wo;
+    int tgsize = M < 64 ? 32 : (M < 256 ? 64 : (M < 512 ? 128 : 256));
+    if (tgsize > 1024) tgsize = 1024;
+
+    ut_tensor* dwp = ut_alloc(1, (int[]){C * l->kh * l->kw}, UT_METAL);
+    _mtl_dispatch_rows(mc, "dwconv2d_bwd_dw", (void*[]){grad_out->gpu_buf, x->gpu_buf, dwp->gpu_buf},
+                       3, params, (int)sizeof(params), C * l->kh * l->kw, tgsize);
+    dwp->dirty_cpu = true;
+    ut_sync_cpu(dwp);
+    ut_sync_cpu(dW);
+    for (int i = 0; i < C * l->kh * l->kw; i++) dW->data[i] += dwp->data[i];
+    dW->dirty_gpu = true;
+    ut_free(dwp);
+
+    if (l->bias && db) {
+      ut_tensor* dbp = ut_alloc(1, (int[]){C}, UT_METAL);
+      _mtl_dispatch_rows(mc, "dwconv2d_bwd_db", (void*[]){grad_out->gpu_buf, dbp->gpu_buf}, 2, params,
+                         (int)sizeof(params), C, tgsize);
+      dbp->dirty_cpu = true;
+      ut_sync_cpu(dbp);
+      ut_sync_cpu(db);
+      for (int c = 0; c < C; c++) db->data[c] += dbp->data[c];
+      db->dirty_gpu = true;
+      ut_free(dbp);
+    }
+
+    ut_tensor* dx = ut_alloc(4, (int[]){N, C, H, W}, UT_METAL);
+    _mtl_dispatch(mc, "dwconv2d_bwd_dx",
+                  (void*[]){grad_out->gpu_buf, l->weight->gpu_buf, dx->gpu_buf}, 3, params,
+                  (int)sizeof(params), N * C * H * W);
+    dx->dirty_cpu = true;
+    return dx;
+  }
+
+  // CPU fallback: scatter-add dx and dW/db together in one pass
+  ut_sync_cpu(grad_out);
+  ut_sync_cpu(x);
+  ut_sync_cpu(l->weight);
+  ut_sync_cpu(dW);
+  if (l->bias && db) ut_sync_cpu(db);
+  ut_tensor* dx = ut_alloc(4, (int[]){N, C, H, W}, UT_CPU);
+  memset(dx->data, 0, (size_t)dx->shape.nelem * sizeof(float));
+  for (int n = 0; n < N; n++)
+    for (int c = 0; c < C; c++)
+      for (int oh = 0; oh < Ho; oh++)
+        for (int ow = 0; ow < Wo; ow++) {
+          float go = grad_out->data[((n * C + c) * Ho + oh) * Wo + ow];
+          if (l->bias && db) db->data[c] += go;
+          for (int khh = 0; khh < l->kh; khh++)
+            for (int kww = 0; kww < l->kw; kww++) {
+              int ih = oh * l->stride - l->pad + khh, iw = ow * l->stride - l->pad + kww;
+              if (ih < 0 || ih >= H || iw < 0 || iw >= W) continue;
+              dW->data[(c * l->kh + khh) * l->kw + kww] +=
+                  go * x->data[((n * C + c) * H + ih) * W + iw];
+              dx->data[((n * C + c) * H + ih) * W + iw] +=
+                  go * l->weight->data[(c * l->kh + khh) * l->kw + kww];
+            }
+        }
+  dW->dirty_gpu = true;
+  if (l->bias && db) db->dirty_gpu = true;
+  dx->dirty_gpu = true;
+  return dx;
+}
+
+void ut_dwconv2d_cache_free(ut_dwconv2d_cache* c) {
+  ut_free(c->input);
+  c->input = NULL;
+}
+void ut_dwconv2d_free(ut_dwconv2d* l) {
+  ut_free_all(l->weight, l->bias);
+  l->weight = l->bias = NULL;
+}
+
 // =========================================================
 // BatchNorm2d
 // =========================================================