← Commits · dd4b9aeb
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
// =========================================================