+ New

utensil

Public
c913eafb757cb4d8b0f36a13b430599f6bddd106
diff --git a/test.c b/test.c
index a8c7362..f400325 100644
--- a/test.c
+++ b/test.c
@@ -263,6 +263,77 @@ static void test_layernorm_backward(ut_dev dev) {
   ut_layernorm_free(&ln);
 }
 
+static void test_im2col(ut_dev dev) {
+  // x[1,1,1,4]=[1,2,3,4], kh=1,kw=3,s=1,p=0 → col[3,2]
+  ut_tensor* x = ut_from_data(4, (int[]){1, 1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
+  ut_tensor* col = ut_im2col(x, 1, 3, 1, 0);
+  ut_sync_cpu(col);
+  assert(col->shape.ndim == 2 && col->shape.shape[0] == 3 && col->shape.shape[1] == 2);
+  assert_data(col, ((float[]){1.f, 2.f, 2.f, 3.f, 3.f, 4.f}), 1e-6f);
+  ut_free(col);
+  ut_free(x);
+}
+
+static void test_col2im(ut_dev dev) {
+  // reconstruct [1,1,1,4] from col above → [1,4,6,4]
+  ut_tensor* col = ut_from_data(2, (int[]){3, 2}, (float[]){1.f, 2.f, 2.f, 3.f, 3.f, 4.f}, dev);
+  ut_tensor* dx = ut_col2im(col, 1, 1, 1, 4, 1, 3, 1, 0);
+  ut_sync_cpu(dx);
+  assert_data(dx, ((float[]){1.f, 4.f, 6.f, 4.f}), 1e-6f);
+  ut_free(dx);
+  ut_free(col);
+}
+
+static void test_conv1d_forward(ut_dev dev) {
+  // in=1,out=1,kw=3,s=1,p=0, no bias
+  // x=[1,2,3,4], W=[1,0,-1] → out=[-2,-2]
+  ut_conv1d l = ut_conv1d_alloc(1, 1, 3, 1, 0, false, dev);
+  ut_free(l.weight);
+  l.weight = ut_from_data(3, (int[]){1, 1, 3}, (float[]){1.f, 0.f, -1.f}, dev);
+
+  ut_tensor* x = ut_from_data(3, (int[]){1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
+  ut_tensor* out = ut_conv1d_forward(&l, x, NULL);
+  ut_sync_cpu(out);
+  assert(out->shape.ndim == 3 && out->shape.shape[0] == 1 && out->shape.shape[1] == 1 &&
+         out->shape.shape[2] == 2);
+  assert_data(out, ((float[]){-2.f, -2.f}), 1e-4f);
+
+  ut_free(out);
+  ut_free(x);
+  ut_conv1d_free(&l);
+}
+
+static void test_conv1d_backward(ut_dev dev) {
+  // x=[1,2,3,4], W=[1,0,-1], go=[1,0]
+  // col = [[1,2],[2,3],[3,4]]  (3×2)
+  // dW = go_t[1,2] @ col^T[2,3] = [1*1+0*2, 1*2+0*3, 1*3+0*4] = [1,2,3]
+  ut_conv1d l = ut_conv1d_alloc(1, 1, 3, 1, 0, false, dev);
+  ut_free(l.weight);
+  l.weight = ut_from_data(3, (int[]){1, 1, 3}, (float[]){1.f, 0.f, -1.f}, dev);
+
+  ut_tensor* x = ut_from_data(3, (int[]){1, 1, 4}, (float[]){1.f, 2.f, 3.f, 4.f}, dev);
+  ut_conv1d_cache c;
+  ut_tensor* out = ut_conv1d_forward(&l, x, &c);
+  ut_free(out);
+
+  ut_tensor* go = ut_from_data(3, (int[]){1, 1, 2}, (float[]){1.f, 0.f}, dev);
+  ut_tensor* dW = ut_alloc(3, (int[]){1, 1, 3}, dev);
+  memset(dW->data, 0, 12);
+
+  ut_tensor* dx = ut_conv1d_backward(&l, &c, go, dW, NULL);
+  ut_sync_cpu(dx);
+  ut_sync_cpu(dW);
+  assert(dx->shape.ndim == 3 && dx->shape.shape[2] == 4);
+  assert_data(dW, ((float[]){1.f, 2.f, 3.f}), 1e-4f);
+
+  ut_free(dx);
+  ut_free(dW);
+  ut_free(go);
+  ut_conv1d_cache_free(&c);
+  ut_free(x);
+  ut_conv1d_free(&l);
+}
+
 static void test_sgd_momentum(void) {
   ut_tensor* p = ut_alloc(1, (int[]){2}, UT_CPU);
   ut_tensor* params[1] = {p};
@@ -314,6 +385,16 @@ int main() {
   test_layernorm_backward(UT_CPU);
   test_layernorm_backward(UT_METAL);
 
+  test_im2col(UT_CPU);
+  test_im2col(UT_METAL);
+  test_col2im(UT_CPU);
+  test_col2im(UT_METAL);
+
+  test_conv1d_forward(UT_CPU);
+  test_conv1d_forward(UT_METAL);
+  test_conv1d_backward(UT_CPU);
+  test_conv1d_backward(UT_METAL);
+
   test_sgd_momentum();
   return 0;
 }
diff --git a/utensil.h b/utensil.h
index bd3c05b..5851acf 100644
--- a/utensil.h
+++ b/utensil.h
@@ -32,7 +32,6 @@ typedef struct ut_linear {
   ut_tensor* weight;  // [in, out]
   ut_tensor* bias;    // [out]
   int nin, nout;
-  bool has_bias;
 } ut_linear;
 
 typedef struct ut_layernorm {
@@ -48,6 +47,17 @@ typedef struct ut_layernorm_cache {
   ut_tensor* rstd;   // per-row reciprocal std-dev
 } ut_layernorm_cache;
 
+typedef struct ut_conv1d {
+  ut_tensor* weight;  // [out_c, in_c, kw]
+  ut_tensor* bias;    // [out_c]
+  int in_c, out_c, kw, stride, pad;
+} ut_conv1d;
+
+typedef struct ut_conv1d_cache {
+  ut_tensor* input;
+  ut_tensor* col;
+} ut_conv1d_cache;
+
 typedef struct ut_sgd {
   ut_tensor** params;    // pointers to model parameters (not owned)
   ut_tensor** grads;     // gradient accumulators (owned)
@@ -135,47 +145,62 @@ static const char* _mtl_src =
     "for(int k=0;k<K;k++)s+=a[ao+k]*b_[bo+k*N];"
     "c[idx]=s;}\n"
     // bias_add: p={N,C,HW}
-    "__kernel void bias_add(__global float*out,__global const float*bias,"
-    "__global const int*p){"
-    "int idx=(int)get_global_id(0);int tot=p[0]*p[1]*p[2];"
-    "if(idx>=tot)return;int c=(idx/p[2])%p[1];out[idx]+=bias[c];}\n"
+    "kernel void bias_add(device float*out,device const float*bias,"
+    "constant int*p,uint idx[[thread_position_in_grid]]){"
+    "int tot=p[0]*p[1]*p[2];"
+    "if((int)idx>=tot)return;"
+    "int c=((int)idx/p[2])%p[1];"
+    "out[idx]+=bias[c];}"
     // relu_bwd
-    "__kernel void relu_bwd(__global const float*go,__global const float*fwd,"
-    "__global float*gi,__global const int*p){"
-    "int idx=(int)get_global_id(0);if(idx<p[0])gi[idx]=fwd[idx]>0.f?go[idx]:0.f;}\n"
+    "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;}"
     // transpose dims 0<->1: p={A,B,inner}
-    "__kernel void transpose_01(__global const float*in,__global float*out,"
-    "__global const int*p){"
-    "int idx=(int)get_global_id(0);"
-    "int A=p[0],B=p[1],inner=p[2];if(idx>=A*B*inner)return;"
-    "int k=idx%inner,t=idx/inner,a=t%A,b=t/A;"
-    "out[idx]=in[(a*B+b)*inner+k];}\n"
+    "kernel void transpose_01("
+    "device const float*in,device float*out,constant int*p,"
+    "uint idx[[thread_position_in_grid]]){"
+    "int A=p[0],B=p[1],inner=p[2];"
+    "if((int)idx>=A*B*inner)return;"
+    "int k=(int)idx%inner,t=(int)idx/inner;"
+    "int a=t%A,b=t/A;"
+    "out[idx]=in[(a*B+b)*inner+k];}"
     // transpose dims 1<->2: p={A,B,C,D}
-    "__kernel void transpose_12(__global const float*i,__global float*o,"
-    "__global const int*p){"
-    "int idx=(int)get_global_id(0);"
-    "int A=p[0],B=p[1],C=p[2],D=p[3];int n=A*B*C*D;if(idx>=n)return;"
+    "kernel void transpose_12(device const float*i,device float*o,"
+    "constant int*p,uint idx[[thread_position_in_grid]]){"
+    "int A=p[0],B=p[1],C=p[2],D=p[3];"
+    "int n=A*B*C*D;if((int)idx>=n)return;"
     "int d_=idx%D,t=idx/D,c=t%C,t2=t/C,b=t2%B,a=t2/B;"
-    "o[a*C*B*D+c*B*D+b*D+d_]=i[a*B*C*D+b*C*D+c*D+d_];}\n"
+    "o[a*C*B*D+c*B*D+b*D+d_]=i[a*B*C*D+b*C*D+c*D+d_];}"
     // row_softmax: one workgroup per row; local sh[] passed as last arg
     // p={outer,C}  global=outer*tgsize  local=tgsize
-    "__kernel void row_softmax("
-    "__global const float*x,__global float*o,__global const int*p,"
-    "__local float* sh){"
-    "int gid=(int)get_group_id(0);int lid=(int)get_local_id(0);"
-    "int tgs=(int)get_local_size(0);int C=p[1];"
-    "__global const float*row=x+gid*C;__global float*orow=o+gid*C;"
+    "kernel void row_softmax("
+    "device const float*x,device float*o,constant int*p,"
+    "uint gid[[threadgroup_position_in_grid]],"
+    "uint lid[[thread_position_in_threadgroup]],"
+    "uint tgs[[threads_per_threadgroup]]){"
+    "int C=p[1];"
+    "device const float*row=x+gid*C;"
+    "device float*orow=o+gid*C;"
+    "threadgroup float shmem[1024];"
+    // phase 1: max reduce
     "float mx=-1e38f;"
-    "for(int i=lid;i<C;i+=tgs){float v=row[i];if(v>mx)mx=v;}"
-    "sh[lid]=mx;barrier(CLK_LOCAL_MEM_FENCE);"
-    "for(int s=tgs/2;s>0;s>>=1){if(lid<s&&sh[lid+s]>sh[lid])sh[lid]=sh[lid+s];"
-    "barrier(CLK_LOCAL_MEM_FENCE);}float gmx=sh[0];"
-    "float loc=0.f;"
-    "for(int i=lid;i<C;i+=tgs){float e=exp(row[i]-gmx);orow[i]=e;loc+=e;}"
-    "sh[lid]=loc;barrier(CLK_LOCAL_MEM_FENCE);"
-    "for(int s=tgs/2;s>0;s>>=1){if(lid<s)sh[lid]+=sh[lid+s];"
-    "barrier(CLK_LOCAL_MEM_FENCE);}float gs_=sh[0];"
-    "for(int i=lid;i<C;i+=tgs)orow[i]/=gs_;}\n"
+    "for(int i=(int)lid;i<C;i+=(int)tgs){"
+    "float v=row[i];if(v>mx)mx=v;}"
+    "shmem[lid]=mx;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint s=tgs/2;s>0;s>>=1){"
+    "if(lid<s&&shmem[lid+s]>shmem[lid])shmem[lid]=shmem[lid+s];"
+    "threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float gmx=shmem[0];"
+    // phase 2: sum of exp
+    "float loc_sum=0;"
+    "for(int i=(int)lid;i<C;i+=(int)tgs){float e=exp(row[i]-gmx);orow[i]=e;loc_sum+=e;}"
+    "shmem[lid]=loc_sum;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint s=tgs/2;s>0;s>>=1){"
+    "if(lid<s)shmem[lid]+=shmem[lid+s];"
+    "threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float gsum=shmem[0];"
+    // phase 3: divide
+    "for(int i=(int)lid;i<C;i+=(int)tgs)orow[i]/=gsum;}"
     // layernorm forward: x[rows, d] -> out[rows, d]
     // also writes xnorm[rows,d], rstd[rows] for backward.
     // one threadgroup per row.  p = {rows, d}.
@@ -190,6 +215,26 @@ static const char* _mtl_src =
     "device const float*row=x+gid*d;"
     "device float*orow=out+gid*d;device float*xnrow=xn+gid*d;"
     "threadgroup float sh[1024];"
+    // mean
+    "float ls=0;for(int i=(int)lid;i<d;i+=(int)tgs)ls+=row[i];"
+    "sh[lid]=ls;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint "
+    "s=tgs/"
+    "2;s>0;s>>=1){if(lid<s)sh[lid]+=sh[lid+s];threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float mu=sh[0]/(float)d;"
+    // var
+    "ls=0;for(int i=(int)lid;i<d;i+=(int)tgs){float v=row[i]-mu;ls+=v*v;}"
+    "sh[lid]=ls;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint "
+    "s=tgs/"
+    "2;s>0;s>>=1){if(lid<s)sh[lid]+=sh[lid+s];threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float rs=rsqrt(sh[0]/(float)d+1e-5f);"
+    "if(lid==0)rstd[gid]=rs;"
+    // output
+    "for(int i=(int)lid;i<d;i+=(int)tgs){"
+    "float xni=(row[i]-mu)*rs;"
+    "xnrow[i]=xni;"
+    "orow[i]=w[i]*xni+b[i];}}"
     // layernorm backward: given grad_out[rows,d], xnorm[rows,d], rstd[rows],
     // weight[d] → dx[rows,d].  dW[d] and db[d] accumulated separately on CPU
     // (only dx needs to be GPU-fast for the training loop hot path).
@@ -206,6 +251,52 @@ static const char* _mtl_src =
     "device float*dxrow=dx+gid*d;"
     "float rs=rs_[gid];"
     "threadgroup float sh[1024];"
+    // sum(go*w)
+    "float s1=0;for(int i=(int)lid;i<d;i+=(int)tgs)s1+=gorow[i]*w[i];"
+    "sh[lid]=s1;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint "
+    "s=tgs/"
+    "2;s>0;s>>=1){if(lid<s)sh[lid]+=sh[lid+s];threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float sum_go_w=sh[0];"
+    // sum(go*w*xn)
+    "s1=0;for(int i=(int)lid;i<d;i+=(int)tgs)s1+=gorow[i]*w[i]*xnrow[i];"
+    "sh[lid]=s1;threadgroup_barrier(mem_flags::mem_threadgroup);"
+    "for(uint "
+    "s=tgs/"
+    "2;s>0;s>>=1){if(lid<s)sh[lid]+=sh[lid+s];threadgroup_barrier(mem_flags::mem_threadgroup);}"
+    "float sum_go_w_xn=sh[0];"
+    // dx
+    "for(int i=(int)lid;i<d;i+=(int)tgs)"
+    "dxrow[i]=rs*(w[i]*gorow[i]-(sum_go_w+xnrow[i]*sum_go_w_xn)/(float)d);}"
+    // col2im (gather, no atomics): col[C*kh*kw,N*Ho*Wo] -> dx[N,C,H,W]
+    // p = {N,C,H,W,kh,kw,stride,pad,Ho,Wo}
+    "kernel void col2im_k("
+    "device const float*col,device float*dx,constant int*p,"
+    "uint idx[[thread_position_in_grid]]){"
+    "int N=p[0],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];"
+    "if((int)idx>=N*C*H*W)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+=col[(c*kh*kw+khh*kw+kww)*(N*Ho*Wo)+n*Ho*Wo+oh*Wo+ow];}"
+    "dx[idx]=sum;}\n"
+    // im2col: x[N,C,H,W] → col[C*kh*kw, N*Ho*Wo]
+    // p = {N,C,H,W,kh,kw,stride,pad,Ho,Wo}
+    "kernel void im2col_k("
+    "device const float*x,device float*col,constant int*p,"
+    "uint idx[[thread_position_in_grid]]){"
+    "int N=p[0],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=C*kh*kw*N*Ho*Wo;if((int)idx>=(int)tot)return;"
+    "int t=(int)idx,ow=t%Wo;t/=Wo;int oh=t%Ho;t/=Ho;"
+    "int kww=t%kw;t/=kw;int khh=t%kh;t/=kh;"
+    "int c=t%C;t/=C;int n=t;"
+    "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"
     "";
 
 #define _MTL_MAX_PL 32  // maximum number of pipeline states
@@ -235,6 +326,7 @@ static void* _mtl_init(void) {
       (id)c->device, sel_getUid("newLibraryWithSource:options:error:"), (id)src, (id)opts, NULL);
   _v0(opts, "release");
   if (!lib) {
+    fprintf(stderr, "utensil: failed to compile Metal library\n");
     _v0(c->queue, "release");
     _v0(c->device, "release");
     free(c);
@@ -780,8 +872,8 @@ static ut_tensor* ut_matmul_t(ut_tensor* a, ut_tensor* b, bool ta, bool tb) {
   int kb = tb ? cb : rb;
   if (k != kb) return NULL;
   ut_dev dev = (a->dev == UT_METAL || b->dev == UT_METAL) ? UT_METAL : UT_CPU;
-  ut_sync_cpu(a);
-  ut_sync_cpu(b);
+  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();
@@ -802,7 +894,7 @@ static ut_tensor* ut_matmul_t(ut_tensor* a, ut_tensor* b, bool ta, bool tb) {
 // Linear layer
 // =========================================================
 ut_linear ut_linear_alloc(int in, int out, bool bias, ut_dev dev) {
-  ut_linear l = {.nin = in, .nout = out, .has_bias = bias};
+  ut_linear l = {.nin = in, .nout = out};
   l.weight = ut_randn(2, (int[]){in, out}, 0.f, sqrtf(2.f / (float)in), dev);
   if (bias) l.bias = ut_alloc(1, (int[]){out}, dev);
   return l;
@@ -812,7 +904,7 @@ ut_linear ut_linear_alloc(int in, int out, bool bias, ut_dev dev) {
 // x:[B,in], W:[in,out] → out:[B,out]
 ut_tensor* ut_linear_forward(ut_linear* l, ut_tensor* x) {
   ut_tensor* out = ut_matmul(x, l->weight);
-  if (l->has_bias) {
+  if (l->bias) {
     _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
     if (out->dev == UT_METAL && mc) {
       ut_sync_gpu(out);
@@ -845,7 +937,7 @@ ut_tensor* ut_linear_backward(ut_linear* l, ut_tensor* x, ut_tensor* grad_out, u
   ut_free(dWb);
 
   // db += sum(grad_out, axis=0)
-  if (l->has_bias && db) {
+  if (l->bias && db) {
     ut_sync_cpu(grad_out);
     ut_sync_cpu(db);
     for (int b = 0; b < B; b++)
@@ -1047,6 +1139,198 @@ void ut_layernorm_free(ut_layernorm* l) {
   l->weight = l->bias = NULL;
 }
 
+// =========================================================
+// im2col / col2im
+// =========================================================
+
+ut_tensor* ut_im2col(ut_tensor* x, int kh, int kw, int stride, int pad) {
+  int N = x->shape.shape[0], C = x->shape.shape[1];
+  int H = x->shape.shape[2], W = x->shape.shape[3];
+  int Ho = (H + 2 * pad - kh) / stride + 1;
+  int Wo = (W + 2 * pad - kw) / stride + 1;
+  int col_dims[2] = {C * kh * kw, N * Ho * Wo};
+
+  _mtl_ctx_t* mc = (_mtl_ctx_t*)ut_metal_ctx();
+  if (mc && x->dev == UT_METAL) {
+    ut_tensor* col = ut_alloc(2, col_dims, UT_METAL);
+    int params[10] = {N, C, H, W, kh, kw, stride, pad, Ho, Wo};
+    void* ib[2] = {x->gpu_buf, col->gpu_buf};
+    _mtl_dispatch(mc, "im2col_k", ib, 2, params, (int)sizeof(params), C * kh * kw * N * Ho * Wo);
+    col->dirty_cpu = true;
+    return col;
+  }
+
+  // CPU path
+  ut_sync_cpu(x);
+  ut_tensor* col = ut_alloc(2, col_dims, UT_CPU);
+  for (int n = 0; n < N; n++)
+    for (int c = 0; c < C; c++)
+      for (int hh = 0; hh < kh; hh++)
+        for (int ww = 0; ww < kw; ww++) {
+          int row = c * kh * kw + hh * kw + ww;
+          for (int oh = 0; oh < Ho; oh++)
+            for (int ow = 0; ow < Wo; ow++) {
+              int ih = oh * stride - pad + hh;
+              int iw = ow * stride - pad + ww;
+              float v = 0;
+              if (ih >= 0 && ih < H && iw >= 0 && iw < W)
+                v = x->data[((n * C + c) * H + ih) * W + iw];
+              col->data[row * (N * Ho * Wo) + n * Ho * Wo + oh * Wo + ow] = v;
+            }
+        }
+  return col;
+}
+
+ut_tensor* ut_col2im(ut_tensor* col, int N, int C, int H, int W, int kh, int kw, int stride,
+                     int pad) {
+  int Ho = (H + 2 * pad - kh) / stride + 1;
+  int Wo = (W + 2 * pad - kw) / stride + 1;
+  int dims[4] = {N, C, H, W};
+
+  _mtl_ctx_t* _mc_c = (_mtl_ctx_t*)ut_metal_ctx();
+  if (_mc_c && col->gpu_buf) {
+    ut_to_device(col, UT_METAL);
+    ut_tensor* out = ut_alloc(4, dims, UT_METAL);
+    int _cp[10] = {N, C, H, W, kh, kw, stride, pad, Ho, Wo};
+    void* _cb[2] = {col->gpu_buf, out->gpu_buf};
+    _mtl_dispatch(_mc_c, "col2im_k", _cb, 2, _cp, (int)sizeof(_cp), N * C * H * W);
+    out->dirty_cpu = true;
+    return out;
+  }
+
+  ut_sync_cpu(col);
+  ut_tensor* out = ut_alloc(4, dims, UT_CPU);
+  for (int n = 0; n < N; n++)
+    for (int c = 0; c < C; c++)
+      for (int hh = 0; hh < kh; hh++)
+        for (int ww = 0; ww < kw; ww++) {
+          int row = c * kh * kw + hh * kw + ww;
+          for (int oh = 0; oh < Ho; oh++)
+            for (int ow = 0; ow < Wo; ow++) {
+              int ih = oh * stride - pad + hh;
+              int iw = ow * stride - pad + ww;
+              if (ih >= 0 && ih < H && iw >= 0 && iw < W)
+                out->data[((n * C + c) * H + ih) * W + iw] +=
+                    col->data[row * (N * Ho * Wo) + n * Ho * Wo + oh * Wo + ow];
+            }
+        }
+  return out;
+}
+
+// =========================================================
+// Conv1D
+// =========================================================
+
+ut_conv1d ut_conv1d_alloc(int in_c, int out_c, int kw, int stride, int pad, bool bias, ut_dev dev) {
+  ut_conv1d l = {.in_c = in_c, .out_c = out_c, .kw = kw, .stride = stride, .pad = pad};
+  float std = sqrtf(2.f / (float)(in_c * kw));
+  l.weight = ut_randn(3, (int[]){out_c, in_c, kw}, 0.f, std, dev);
+  if (bias) l.bias = ut_alloc(1, (int[]){out_c}, dev);
+  return l;
+}
+
+// x: [N, C, L] — treat as [N, C, 1, L] for conv2d im2col
+ut_tensor* ut_conv1d_forward(ut_conv1d* l, ut_tensor* x, ut_conv1d_cache* cache) {
+  int N = x->shape.shape[0], L = x->shape.shape[2];
+  int Lo = (L + 2 * l->pad - l->kw) / l->stride + 1;
+  // expand x to [N, C, 1, L]
+  int xd4[4] = {N, l->in_c, 1, L};
+  ut_reshape(x, 4, xd4);  // temporarily treat as [N, in_c, 1, L]
+
+  // Use im2col with kh=1, stride_h=1, pad_h=0
+  ut_tensor* col = ut_im2col(x, 1, l->kw, l->stride, l->pad);
+
+  int w2d[2] = {l->out_c, l->in_c * l->kw};
+  ut_tensor w_view = *l->weight;
+  w_view.shape = ut_shape_new(2, w2d);
+  w_view.rc = 0x7fffffff;
+
+  ut_tensor* out2 = ut_matmul(&w_view, col);  // [out_c, N*Lo]
+  ut_sync_cpu(out2);
+  // out2 is [out_c, N*Lo] — reshape to [out_c, N, Lo] then transpose(0,1)
+  int cn[3] = {l->out_c, N, Lo};
+  ut_reshape(out2, 3, cn);
+  ut_tensor* out = ut_transpose(out2, 0, 1);  // [N, out_c, Lo]
+  ut_free(out2);
+
+  if (l->bias) {
+    ut_sync_cpu(out);
+    ut_sync_cpu(l->bias);
+    for (int n = 0; n < N; n++)
+      for (int oc = 0; oc < l->out_c; oc++)
+        for (int lp = 0; lp < Lo; lp++)
+          out->data[(n * l->out_c + oc) * Lo + lp] += l->bias->data[oc];
+  }
+  // restore x to original 3D shape before cache retain
+  int x3d[3] = {N, l->in_c, L};
+  ut_reshape(x, 3, x3d);
+  if (cache) {
+    cache->input = ut_retain(x);
+    cache->col = ut_retain(col);
+  }
+  ut_free(col);
+  return out;
+}
+
+ut_tensor* ut_conv1d_backward(ut_conv1d* l, ut_conv1d_cache* cache, ut_tensor* grad_out,
+                              ut_tensor* dW, ut_tensor* db) {
+  ut_tensor* col = cache->col;
+  int N = cache->input->shape.shape[0];
+  int L = cache->input->shape.shape[2];
+  int Lo = (L + 2 * l->pad - l->kw) / l->stride + 1;
+  ut_sync_cpu(grad_out);
+  ut_sync_cpu(col);
+  ut_sync_cpu(dW);
+
+  // go_2d: [out_c, N*Lo] — transpose + reshape grad_out [N,out_c,Lo]
+  ut_tensor* go_t = ut_transpose(grad_out, 0, 1);  // [out_c, N, Lo]
+  int d2[2] = {l->out_c, N * Lo};
+  ut_reshape(go_t, 2, d2);  // [out_c, N*Lo]
+
+  _mtl_ctx_t* _mc = (_mtl_ctx_t*)ut_metal_ctx();
+  if (_mc) ut_to_device(go_t, UT_METAL);
+
+  ut_tensor* dWb = ut_matmul_t(go_t, col, false, true);
+  ut_sync_cpu(dWb);
+  ut_sync_cpu(dW);
+  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 lp = 0; lp < Lo; lp++)
+          db->data[oc] += grad_out->data[(n * l->out_c + oc) * Lo + lp];
+  }
+
+  int w2d[2] = {l->out_c, l->in_c * l->kw};
+  ut_tensor w_view = *l->weight;
+  w_view.shape = ut_shape_new(2, w2d);
+  w_view.rc = 0x7fffffff;
+  ut_tensor* dcol = ut_matmul_t(&w_view, go_t, true, false);
+  ut_free(go_t);
+
+  // col2im back to [N, C, 1, L] then squeeze to [N, C, L]
+  ut_tensor* dx4 = ut_col2im(dcol, N, l->in_c, 1, L, 1, l->kw, l->stride, l->pad);
+  ut_free(dcol);
+  int dx3d[3] = {N, l->in_c, L};
+  ut_reshape(dx4, 3, dx3d);
+  return dx4;
+}
+
+void ut_conv1d_cache_free(ut_conv1d_cache* c) {
+  ut_free(c->input);
+  ut_free(c->col);
+  c->input = c->col = NULL;
+}
+void ut_conv1d_free(ut_conv1d* l) {
+  ut_free(l->weight);
+  if (l->bias) ut_free(l->bias);
+  l->weight = l->bias = NULL;
+}
+
 // =========================================================
 // Softmax
 // =========================================================