code wiki / (root) / nx_vitpose_forward.nx

nx_vitpose_forward.nx source

↩ module page · 178 lines · 11190 B

1// nx_vitpose_forward.nx -- THE CULMINATING RUNG: the full faithful f32 ViTPose-base forward, composing every gated 2// rung on the REAL downloaded weights. Reads the 343MB safetensors into RAM once, loads tensors by name (per-layer 3// buffers reused across the 12 layers so only ~57MB of expanded weights live at a time), then: 4// patch-embed conv -> 192 tokens + pos_embed[:,1:] -> 12x nx_vit_encoder_layer -> final backbone.layernorm -> 5// reshape 192->768x16x12 -> ReLU -> bilinear x4 -> head.conv(768->17,k3,p1) -> ReLU -> nx_pose_extract -> 17 kpts. 6// (ReLU before the integer keypoint argmax: the trained heatmap peak is positive, so it preserves the peak cell.) 7// Structural validation on a synthetic image: runs end-to-end, 17 finite heatmaps, 17 keypoints in the 48x64 (x4) 8// quarter-pixel range. Semantic validation (keypoints on a real person) = the follow-on image pipeline. expect_exit: 0 9import "nx_syscalls.nx" 10import "nx_f32.nx" 11import "nx_f32_cvt.nx" 12import "nx_f32_conv2d.nx" 13import "nx_f32_layernorm.nx" 14import "nx_f32_upsample_bilinear.nx" 15import "nx_pose_cnn.nx" 16import "nx_pose_keypoints.nx" 17import "nx_safetensors_load.nx" 18import "nx_vit_encoder_layer.nx" 19const K_MAGIC_4096: i64 = 4096 20const K_MAGIC_345000000: i64 = 345000000 21const K_MAGIC_3072: i64 = 3072 22const K_MAGIC_1000000000: i64 = 1000000000 23 24func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 25func wn(v: i64) -> i64 { var m: i64=v; if m<0{w("-" as *u8);m=0-m} let t:*u8=sys_mmap(24); var k:i64=0; if m==0{t[0]=48 as u8;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i:i64=0; let o:*u8=sys_mmap(24); while i<k{o[i]=t[k-1-i];i=i+1} sys_write(1,o,k); return 0 } 26func rdcfg(path: *u8, buf: *u8, cap: i64) -> i64 { let fd: i64=sys_openat_rd(path); if fd<0{return 0-1} var n: i64=sys_read(fd,buf,cap-1); sys_close(fd); if n<0{return 0-1} var go: i64=1; while go==1{go=0; if n>0{ let c: i64=buf[n-1] as i64; if c==10{n=n-1;go=1} if c==13{n=n-1;go=1} if c==32{n=n-1;go=1} }} buf[n]=0 as u8; return n } 27func u64le(buf: *u8, off: i64) -> i64 { var v: i64=0; var i: i64=0; while i<8 { v=v|((buf[off+i]&0xff)<<(i*8)); i=i+1 } return v } 28func catn(dst: *u8, off: i64, s: *u8) -> i64 { var o: i64=off; var i: i64=0; while s[i]!=(0 as u8){dst[o]=s[i];o=o+1;i=i+1} return o } 29func catnum(dst: *u8, off: i64, v: i64) -> i64 { var m: i64=v; let t: *u8=sys_mmap(24); var k: i64=0; if m==0{t[0]=48 as u8;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var o: i64=off; var i: i64=0; while i<k{dst[o]=t[k-1-i];o=o+1;i=i+1} return o } 30 31// load a named tensor from the in-RAM file buffer into out. Returns element count or -1. 32func loadt(fb: *u8, hlen: i64, ds: i64, name: *u8, out: *i64) -> i64 { 33 let dtype: *u8=sys_mmap(16); let offs: *i64=sys_mmap(16) as *i64 34 if stl_tensor(fb, hlen, name, dtype, offs) == 0 { return 0-1 } 35 return stl_read_f32(fb, ds, offs, dtype, out) 36} 37 38func main() -> i64 { 39 // ---- read the whole safetensors into RAM ---- 40 let path: *u8 = sys_mmap(K_MAGIC_4096) 41 if rdcfg("data/mp_fetch_dest.txt" as *u8, path, K_MAGIC_4096) < 0 { w("no dest\n" as *u8); sys_exit(1); return 1 } 42 let fd: i64 = sys_openat_rd(path); if fd<0 { w("open fail\n" as *u8); sys_exit(1); return 1 } 43 let CAP: i64 = K_MAGIC_345000000 44 let fb: *u8 = sys_mmap(CAP) 45 var got: i64 = 0 46 var go: i64 = 1 47 while go == 1 { let n: i64 = sys_read(fd, ((fb as i64)+got) as *u8, CAP-got); if n<=0 { go=0 } else { got=got+n } } 48 sys_close(fd) 49 let hlen: i64 = u64le(fb, 0); let ds: i64 = 8 + hlen 50 w("loaded file bytes=" as *u8); wn(got); w(" header=" as *u8); wn(hlen); w("\n" as *u8) 51 52 // ---- patch embed ---- 53 let pw: *i64 = sys_mmap(8*768*3*16*16) as *i64 54 let pb: *i64 = sys_mmap(8*768) as *i64 55 let pe: *i64 = sys_mmap(8*193*768) as *i64 56 loadt(fb, hlen, ds, "\"backbone.embeddings.patch_embeddings.projection.weight\":" as *u8, pw) 57 loadt(fb, hlen, ds, "\"backbone.embeddings.patch_embeddings.projection.bias\":" as *u8, pb) 58 loadt(fb, hlen, ds, "\"backbone.embeddings.position_embeddings\":" as *u8, pe) 59 let img: *i64 = sys_mmap(8*3*256*192) as *i64 60 let iN: i64 = 3*256*192 61 var i: i64 = 0 62 let ifd: i64 = sys_openat_rd("/home/elderwesto/vitpose_input.bin" as *u8) 63 if ifd >= 0 { 64 let ib: *u8 = sys_mmap(iN*4 + 16) 65 var ig: i64 = 0 66 while ig < iN*4 { let rn: i64 = sys_read(ifd, ((ib as i64)+ig) as *u8, iN*4-ig); if rn<=0{ig=iN*4}else{ig=ig+rn} } 67 sys_close(ifd) 68 i = 0 69 while i < iN { img[i] = (ib[i*4]&0xff)|((ib[i*4+1]&0xff)<<8)|((ib[i*4+2]&0xff)<<16)|((ib[i*4+3]&0xff)<<24); i=i+1 } 70 w("using REAL image input (vitpose_input.bin)\n" as *u8) 71 } else { 72 i = 0; while i < iN { img[i] = nx_i32_to_f32((i % 7) - 3); i = i + 1 } 73 w("using synthetic input\n" as *u8) 74 } 75 let feat0: *i64 = sys_mmap(8*768*16*12) as *i64 76 nx_f32_conv2d_forward(img, 1, 3, 256, 192, pw, 768, 16, 16, 16, 0, pb, feat0) 77 let tok: *i64 = sys_mmap(8*192*768) as *i64 78 var p: i64 = 0 79 while p < 192 { var c: i64 = 0; while c < 768 { tok[p*768+c] = nx_f32_add(feat0[c*192+p], pe[(1+p)*768+c]); c=c+1 } p=p+1 } 80 w("patch-embed -> 192x768 tokens\n" as *u8) 81 82 // ---- 12 encoder layers (reused weight buffers) ---- 83 let Wq: *i64=sys_mmap(8*768*768) as *i64; let bq: *i64=sys_mmap(8*768) as *i64 84 let Wk: *i64=sys_mmap(8*768*768) as *i64; let bk: *i64=sys_mmap(8*768) as *i64 85 let Wv: *i64=sys_mmap(8*768*768) as *i64; let bv: *i64=sys_mmap(8*768) as *i64 86 let Wo: *i64=sys_mmap(8*768*768) as *i64; let bo: *i64=sys_mmap(8*768) as *i64 87 let l1g: *i64=sys_mmap(8*768) as *i64; let l1b: *i64=sys_mmap(8*768) as *i64 88 let l2g: *i64=sys_mmap(8*768) as *i64; let l2b: *i64=sys_mmap(8*768) as *i64 89 let f1w: *i64=sys_mmap(8*K_MAGIC_3072*768) as *i64; let f1b: *i64=sys_mmap(8*K_MAGIC_3072) as *i64 90 let f2w: *i64=sys_mmap(8*768*K_MAGIC_3072) as *i64; let f2b: *i64=sys_mmap(8*768) as *i64 91 let wts: *i64=sys_mmap(8*16) as *i64 92 wts[0]=l1g as i64; wts[1]=l1b as i64; wts[2]=Wq as i64; wts[3]=bq as i64; wts[4]=Wk as i64; wts[5]=bk as i64 93 wts[6]=Wv as i64; wts[7]=bv as i64; wts[8]=Wo as i64; wts[9]=bo as i64; wts[10]=l2g as i64; wts[11]=l2b as i64 94 wts[12]=f1w as i64; wts[13]=f1b as i64; wts[14]=f2w as i64; wts[15]=f2b as i64 95 let bufs: *i64 = sys_mmap(8*16) as *i64 96 bufs[0]=l1g as i64; bufs[1]=l1b as i64; bufs[2]=Wq as i64; bufs[3]=bq as i64; bufs[4]=Wk as i64; bufs[5]=bk as i64 97 bufs[6]=Wv as i64; bufs[7]=bv as i64; bufs[8]=Wo as i64; bufs[9]=bo as i64; bufs[10]=l2g as i64; bufs[11]=l2b as i64 98 bufs[12]=f1w as i64; bufs[13]=f1b as i64; bufs[14]=f2w as i64; bufs[15]=f2b as i64 99 let sfx: *i64 = sys_mmap(8*16) as *i64 100 sfx[0]="layernorm_before.weight\":" as *u8 as i64; sfx[1]="layernorm_before.bias\":" as *u8 as i64 101 sfx[2]="attention.attention.query.weight\":" as *u8 as i64; sfx[3]="attention.attention.query.bias\":" as *u8 as i64 102 sfx[4]="attention.attention.key.weight\":" as *u8 as i64; sfx[5]="attention.attention.key.bias\":" as *u8 as i64 103 sfx[6]="attention.attention.value.weight\":" as *u8 as i64; sfx[7]="attention.attention.value.bias\":" as *u8 as i64 104 sfx[8]="attention.output.dense.weight\":" as *u8 as i64; sfx[9]="attention.output.dense.bias\":" as *u8 as i64 105 sfx[10]="layernorm_after.weight\":" as *u8 as i64; sfx[11]="layernorm_after.bias\":" as *u8 as i64 106 sfx[12]="mlp.fc1.weight\":" as *u8 as i64; sfx[13]="mlp.fc1.bias\":" as *u8 as i64 107 sfx[14]="mlp.fc2.weight\":" as *u8 as i64; sfx[15]="mlp.fc2.bias\":" as *u8 as i64 108 109 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_1000000000)) // ~1e-9 (HF 1e-12; negligible either way) 110 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(64))) // 1/8 111 let nm: *u8 = sys_mmap(256) 112 let outtok: *i64 = sys_mmap(8*192*768) as *i64 113 var L: i64 = 0 114 while L < 12 { 115 var t: i64 = 0 116 while t < 16 { 117 var o: i64 = catn(nm, 0, "\"backbone.encoder.layer." as *u8) 118 o = catnum(nm, o, L); o = catn(nm, o, "." as *u8); o = catn(nm, o, (sfx[t]) as *u8); nm[o]=0 as u8 119 loadt(fb, hlen, ds, nm, (bufs[t]) as *i64) 120 t = t + 1 121 } 122 vit_encoder_layer(tok, wts, 192, 768, 12, 64, K_MAGIC_3072, eps, scale, outtok) 123 var k: i64 = 0; while k < 192*768 { tok[k] = outtok[k]; k = k + 1 } 124 w(" layer " as *u8); wn(L); w(" done\n" as *u8) 125 L = L + 1 126 } 127 128 // ---- final backbone layernorm ---- 129 loadt(fb, hlen, ds, "\"backbone.layernorm.weight\":" as *u8, l1g) 130 loadt(fb, hlen, ds, "\"backbone.layernorm.bias\":" as *u8, l1b) 131 nx_f32_layernorm(tok, l1g, l1b, 192, 768, eps, outtok) 132 133 // ---- decoder: reshape -> ReLU -> bilinear x4 -> head.conv -> ReLU -> keypoints ---- 134 let feat: *i64 = sys_mmap(8*768*16*12) as *i64 135 p = 0; while p < 192 { var c: i64 = 0; while c < 768 { feat[c*192+p] = outtok[p*768+c]; c=c+1 } p=p+1 } 136 pose_cnn_relu(feat, 768*16*12) 137 let up: *i64 = sys_mmap(8*768*64*48) as *i64 138 nx_f32_upsample_bilinear(feat, 1, 768, 16, 12, 4, up) 139 let hcw: *i64 = sys_mmap(8*17*768*3*3) as *i64 140 let hcb: *i64 = sys_mmap(8*17) as *i64 141 loadt(fb, hlen, ds, "\"head.conv.weight\":" as *u8, hcw) 142 loadt(fb, hlen, ds, "\"head.conv.bias\":" as *u8, hcb) 143 let heat: *i64 = sys_mmap(8*17*64*48) as *i64 144 nx_f32_conv2d_forward(up, 1, 768, 64, 48, hcw, 17, 3, 3, 1, 1, hcb, heat) 145 w("decoder -> 17 heatmaps 64x48\n" as *u8) 146 147 // ---- keypoints: f32-aware argmax per heatmap (TRUE max, handles negatives -- no ReLU-collapse hack) ---- 148 let qx: *i64=sys_mmap(8*17) as *i64; let qy: *i64=sys_mmap(8*17) as *i64 149 var jj: i64 = 0 150 while jj < 17 { 151 let hm: i64 = jj*64*48 152 var best: i64 = heat[hm]; var bx: i64 = 0; var by: i64 = 0 153 var yy: i64 = 0 154 while yy < 64 { 155 var xx: i64 = 0 156 while xx < 48 { 157 let v: i64 = heat[hm + yy*48 + xx] 158 if nx_f32_gt(v, best) == 1 { best = v; bx = xx; by = yy } 159 xx = xx + 1 160 } 161 yy = yy + 1 162 } 163 qx[jj] = bx; qy[jj] = by 164 jj = jj + 1 165 } 166 var bad: i64=0; i=0; while i<17*64*48 { if nx_f32_is_nan(heat[i])==1{bad=bad+1} if nx_f32_is_inf(heat[i])==1{bad=bad+1} i=i+1 } 167 var minx: i64=48; var maxx: i64=0; var miny: i64=64; var maxy: i64=0 168 i=0; while i<17 { if qx[i]<minx{minx=qx[i]} if qx[i]>maxx{maxx=qx[i]} if qy[i]<miny{miny=qy[i]} if qy[i]>maxy{maxy=qy[i]} i=i+1 } 169 w("keypoints: 17 (pixel W48xH64) heatmap NaN/Inf=" as *u8); wn(bad); w("\n" as *u8) 170 w(" SPREAD bbox: x[" as *u8); wn(minx); w(".." as *u8); wn(maxx); w("] y[" as *u8); wn(miny); w(".." as *u8); wn(maxy); w("] = " as *u8); wn(maxx-minx); w("x" as *u8); wn(maxy-miny); w("\n" as *u8) 171 w(" all: " as *u8) 172 i=0; while i<17 { w("(" as *u8); wn(qx[i]); w("," as *u8); wn(qy[i]); w(")" as *u8); i=i+1 } w("\n" as *u8) 173 174 if bad==0 { w("VITPOSE-FORWARD OK: full faithful f32 ViTPose ran end-to-end on real image -> 17 keypoints\n" as *u8); sys_exit(0); return 0 } 175 w("VITPOSE-FORWARD FAIL\n" as *u8) 176 sys_exit(1) 177 return 1 178}