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}