code wiki / (root) / nx_vit_encoder_layer.nx

nx_vit_encoder_layer.nx source

↩ module page · 160 lines · 7792 B

1// nx_vit_encoder_layer.nx -- the faithful f32 PRE-NORM ViT encoder block (the ViTPose repeated unit), composing the 2// gated blocks: h = x + attn_dense(MHSA(LN_before(x))); out = h + fc2(GELU(fc1(LN_after(h)))). MHSA has explicit 3// qkv_bias and HF [out,in] weights (via nx_f32_matmul_t = x@W^T), 12 heads x head_dim 64, scale 1/sqrt(head_dim), 4// non-causal. Weights are passed as a 16-pointer array to stay under the arg cap. license_tier: ORIGINAL 5import "nx_syscalls.nx" 6import "nx_f32.nx" 7import "nx_f32_matmul_t.nx" 8import "nx_f32_softmax.nx" 9import "nx_f32_layernorm.nx" 10import "nx_f32_gelu.nx" 11import "nx_thread_pool.nx" // vit_mhsa_pool: per-head attention tasks + pooled projections (2026-09-14) 12 13// add bias[dim] to each of n_rows rows of x, in place. 14func vit_bias_add(x: *i64, bias: *i64, n_rows: i64, dim: i64) -> i64 { 15 var r: i64 = 0 16 while r < n_rows { var c: i64 = 0; while c < dim { x[r*dim + c] = __f32_add(x[r*dim + c], bias[c]); c = c + 1 } r = r + 1 } 17 return 0 18} 19 20// ---- ONE per-head attention kernel shared by the serial vit_mhsa and the pooled vit_mhsa_pool (2026-09-14, search R0 21// cross-encoder). Inner loops on the __f32 hardware intrinsics: IEEE round-to-nearest either way, so the output is 22// BIT-IDENTICAL to the software nx_f32 path it replaces (the argument nx_f32_exp recorded on 2026-07-10, one instruction 23// against the software brick's ~30); nx_vit_encoder_layer_gate is the witness. A pool task per head changes only which 24// thread runs which head, never the accumulation order, so pooled == serial bit for bit. 25struct NxVitAttnCtx { 26 q_ptr: i64, 27 k_ptr: i64, 28 v_ptr: i64, 29 ctx_ptr: i64, 30 t_dim: i64, 31 d_dim: i64, 32 hd_dim: i64, 33 scale: i64, 34 head: i64, 35} 36const NX_VIT_ATTN_CTX_BYTES: i64 = 72 37 38// ctx[i, ho..ho+hd) = softmax(Q_i K^T * scale) V for head h over all T rows 39func vit_attn_head(Q: *i64, K: *i64, V: *i64, ctx: *i64, T: i64, D: i64, hd: i64, scale: i64, h: i64) -> i64 { 40 let sc: *i64 = sys_mmap(8 * T) as *i64 41 let sm: *i64 = sys_mmap(8 * T) as *i64 42 let ho: i64 = h * hd 43 var i: i64 = 0 44 while i < T { 45 var j: i64 = 0 46 while j < T { 47 var s: i64 = 0 48 var d: i64 = 0 49 while d < hd { s = __f32_add(s, __f32_mul(Q[i*D + ho + d], K[j*D + ho + d])); d = d + 1 } 50 sc[j] = __f32_mul(s, scale) 51 j = j + 1 52 } 53 nx_f32_softmax(sc, T, sm) 54 var d2: i64 = 0 55 while d2 < hd { 56 var acc: i64 = 0 57 var jj: i64 = 0 58 while jj < T { acc = __f32_add(acc, __f32_mul(sm[jj], V[jj*D + ho + d2])); jj = jj + 1 } 59 ctx[i*D + ho + d2] = acc 60 d2 = d2 + 1 61 } 62 i = i + 1 63 } 64 sys_munmap(sc as *u8, 8 * T); sys_munmap(sm as *u8, 8 * T) 65 return 0 66} 67 68func _vit_attn_task(ctx_i: i64) -> i64 { 69 let cx: *NxVitAttnCtx = ctx_i as *NxVitAttnCtx 70 return vit_attn_head(cx.q_ptr as *i64, cx.k_ptr as *i64, cx.v_ptr as *i64, cx.ctx_ptr as *i64, cx.t_dim, cx.d_dim, cx.hd_dim, cx.scale, cx.head) 71} 72 73// multi-head self-attention with qkv_bias + output dense. x,out are [T,D]. Wq/Wk/Wv/Wo are [D,D] ([out,in] HF). 74func vit_mhsa(x: *i64, Wq: *i64, bq: *i64, Wk: *i64, bk: *i64, Wv: *i64, bv: *i64, Wo: *i64, bo: *i64, 75 T: i64, D: i64, nh: i64, hd: i64, scale: i64, out: *i64) -> i64 { 76 let Q: *i64 = sys_mmap(8 * T * D) as *i64 77 let K: *i64 = sys_mmap(8 * T * D) as *i64 78 let V: *i64 = sys_mmap(8 * T * D) as *i64 79 nx_f32_matmul_t_blocked(x, Wq, Q, T, D, D); vit_bias_add(Q, bq, T, D) 80 nx_f32_matmul_t_blocked(x, Wk, K, T, D, D); vit_bias_add(K, bk, T, D) 81 nx_f32_matmul_t_blocked(x, Wv, V, T, D, D); vit_bias_add(V, bv, T, D) 82 let ctx: *i64 = sys_mmap(8 * T * D) as *i64 83 var h: i64 = 0 84 while h < nh { vit_attn_head(Q, K, V, ctx, T, D, hd, scale, h); h = h + 1 } 85 nx_f32_matmul_t_blocked(ctx, Wo, out, T, D, D); vit_bias_add(out, bo, T, D) 86 // free the per-call scratch (else 12 layers accumulate ~35MB each -> OOM mid-forward) 87 sys_munmap(Q as *u8, 8*T*D); sys_munmap(K as *u8, 8*T*D); sys_munmap(V as *u8, 8*T*D) 88 sys_munmap(ctx as *u8, 8*T*D) 89 return 0 90} 91 92// pooled twin of vit_mhsa: Q/K/V/O projections through nx_f32_matmul_t_pool and one task per head on the caller's 93// pool. Same signature with the pool in front; bit-identical output (see vit_attn_head). Returns 0, or 1 when the 94// pool wait reports a failure. 95func vit_mhsa_pool(pool: *NxThreadPool, x: *i64, Wq: *i64, bq: *i64, Wk: *i64, bk: *i64, Wv: *i64, bv: *i64, Wo: *i64, bo: *i64, 96 T: i64, D: i64, nh: i64, hd: i64, scale: i64, out: *i64) -> i64 { 97 let Q: *i64 = sys_mmap(8 * T * D) as *i64 98 let K: *i64 = sys_mmap(8 * T * D) as *i64 99 let V: *i64 = sys_mmap(8 * T * D) as *i64 100 nx_f32_matmul_t_pool(pool, x, Wq, Q, T, D, D); vit_bias_add(Q, bq, T, D) 101 nx_f32_matmul_t_pool(pool, x, Wk, K, T, D, D); vit_bias_add(K, bk, T, D) 102 nx_f32_matmul_t_pool(pool, x, Wv, V, T, D, D); vit_bias_add(V, bv, T, D) 103 let ctx: *i64 = sys_mmap(8 * T * D) as *i64 104 let cxs: *u8 = sys_mmap(nh * NX_VIT_ATTN_CTX_BYTES) 105 let done0: i64 = nx_pool_n_completed(pool) 106 var h: i64 = 0 107 while h < nh { 108 let cx: *NxVitAttnCtx = ((cxs as i64) + h * NX_VIT_ATTN_CTX_BYTES) as *NxVitAttnCtx 109 cx.q_ptr = Q as i64 110 cx.k_ptr = K as i64 111 cx.v_ptr = V as i64 112 cx.ctx_ptr = ctx as i64 113 cx.t_dim = T 114 cx.d_dim = D 115 cx.hd_dim = hd 116 cx.scale = scale 117 cx.head = h 118 nx_pool_submit(pool, _vit_attn_task, cx as i64) 119 h = h + 1 120 } 121 let wv: i64 = nx_pool_wait(pool, done0 + nh) 122 sys_munmap(cxs, nh * NX_VIT_ATTN_CTX_BYTES) 123 nx_f32_matmul_t_pool(pool, ctx, Wo, out, T, D, D); vit_bias_add(out, bo, T, D) 124 sys_munmap(Q as *u8, 8*T*D); sys_munmap(K as *u8, 8*T*D); sys_munmap(V as *u8, 8*T*D) 125 sys_munmap(ctx as *u8, 8*T*D) 126 if wv != 0 { return 1 } 127 return 0 128} 129 130// wts[16]: 0 ln1_g,1 ln1_b,2 Wq,3 bq,4 Wk,5 bk,6 Wv,7 bv,8 Wo,9 bo,10 ln2_g,11 ln2_b,12 fc1_w,13 fc1_b,14 fc2_w,15 fc2_b 131func vit_encoder_layer(x: *i64, wts: *i64, T: i64, D: i64, nh: i64, hd: i64, mlp: i64, eps: i64, scale: i64, out: *i64) -> i64 { 132 let ln1g: *i64 = (wts[0]) as *i64; let ln1b: *i64 = (wts[1]) as *i64 133 let Wq: *i64 = (wts[2]) as *i64; let bq: *i64 = (wts[3]) as *i64 134 let Wk: *i64 = (wts[4]) as *i64; let bk: *i64 = (wts[5]) as *i64 135 let Wv: *i64 = (wts[6]) as *i64; let bv: *i64 = (wts[7]) as *i64 136 let Wo: *i64 = (wts[8]) as *i64; let bo: *i64 = (wts[9]) as *i64 137 let ln2g: *i64 = (wts[10]) as *i64; let ln2b: *i64 = (wts[11]) as *i64 138 let fc1w: *i64 = (wts[12]) as *i64; let fc1b: *i64 = (wts[13]) as *i64 139 let fc2w: *i64 = (wts[14]) as *i64; let fc2b: *i64 = (wts[15]) as *i64 140 141 let xn: *i64 = sys_mmap(8 * T * D) as *i64 142 nx_f32_layernorm(x, ln1g, ln1b, T, D, eps, xn) 143 let attn: *i64 = sys_mmap(8 * T * D) as *i64 144 vit_mhsa(xn, Wq, bq, Wk, bk, Wv, bv, Wo, bo, T, D, nh, hd, scale, attn) 145 let hbuf: *i64 = sys_mmap(8 * T * D) as *i64 146 var i: i64 = 0 147 while i < T*D { hbuf[i] = nx_f32_add(x[i], attn[i]); i = i + 1 } 148 let hn: *i64 = sys_mmap(8 * T * D) as *i64 149 nx_f32_layernorm(hbuf, ln2g, ln2b, T, D, eps, hn) 150 let f1: *i64 = sys_mmap(8 * T * mlp) as *i64 151 nx_f32_matmul_t(hn, fc1w, f1, T, D, mlp); vit_bias_add(f1, fc1b, T, mlp) 152 nx_f32_gelu_vec(f1, T * mlp) 153 let m2: *i64 = sys_mmap(8 * T * D) as *i64 154 nx_f32_matmul_t(f1, fc2w, m2, T, mlp, D); vit_bias_add(m2, fc2b, T, D) 155 i = 0 156 while i < T*D { out[i] = nx_f32_add(hbuf[i], m2[i]); i = i + 1 } 157 sys_munmap(xn as *u8, 8*T*D); sys_munmap(attn as *u8, 8*T*D); sys_munmap(hbuf as *u8, 8*T*D) 158 sys_munmap(hn as *u8, 8*T*D); sys_munmap(f1 as *u8, 8*T*mlp); sys_munmap(m2 as *u8, 8*T*D) 159 return 0 160}