code wiki / (root) / nx_vit_encoder_layer.nx

nx_vit_encoder_layer.nx source

↩ module page · 94 lines · 4701 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" 11 12// add bias[dim] to each of n_rows rows of x, in place. 13func vit_bias_add(x: *i64, bias: *i64, n_rows: i64, dim: i64) -> i64 { 14 var r: i64 = 0 15 while r < n_rows { var c: i64 = 0; while c < dim { x[r*dim + c] = nx_f32_add(x[r*dim + c], bias[c]); c = c + 1 } r = r + 1 } 16 return 0 17} 18 19// multi-head self-attention with qkv_bias + output dense. x,out are [T,D]. Wq/Wk/Wv/Wo are [D,D] ([out,in] HF). 20func vit_mhsa(x: *i64, Wq: *i64, bq: *i64, Wk: *i64, bk: *i64, Wv: *i64, bv: *i64, Wo: *i64, bo: *i64, 21 T: i64, D: i64, nh: i64, hd: i64, scale: i64, out: *i64) -> i64 { 22 let Q: *i64 = sys_mmap(8 * T * D) as *i64 23 let K: *i64 = sys_mmap(8 * T * D) as *i64 24 let V: *i64 = sys_mmap(8 * T * D) as *i64 25 nx_f32_matmul_t(x, Wq, Q, T, D, D); vit_bias_add(Q, bq, T, D) 26 nx_f32_matmul_t(x, Wk, K, T, D, D); vit_bias_add(K, bk, T, D) 27 nx_f32_matmul_t(x, Wv, V, T, D, D); vit_bias_add(V, bv, T, D) 28 let ctx: *i64 = sys_mmap(8 * T * D) as *i64 29 let sc: *i64 = sys_mmap(8 * T) as *i64 30 let sm: *i64 = sys_mmap(8 * T) as *i64 31 var h: i64 = 0 32 while h < nh { 33 let ho: i64 = h * hd 34 var i: i64 = 0 35 while i < T { 36 var j: i64 = 0 37 while j < T { 38 var s: i64 = 0 39 var d: i64 = 0 40 while d < hd { s = nx_f32_add(s, nx_f32_mul(Q[i*D + ho + d], K[j*D + ho + d])); d = d + 1 } 41 sc[j] = nx_f32_mul(s, scale) 42 j = j + 1 43 } 44 nx_f32_softmax(sc, T, sm) 45 var d2: i64 = 0 46 while d2 < hd { 47 var acc: i64 = 0 48 var jj: i64 = 0 49 while jj < T { acc = nx_f32_add(acc, nx_f32_mul(sm[jj], V[jj*D + ho + d2])); jj = jj + 1 } 50 ctx[i*D + ho + d2] = acc 51 d2 = d2 + 1 52 } 53 i = i + 1 54 } 55 h = h + 1 56 } 57 nx_f32_matmul_t(ctx, Wo, out, T, D, D); vit_bias_add(out, bo, T, D) 58 // free the per-call scratch (else 12 layers accumulate ~35MB each -> OOM mid-forward) 59 sys_munmap(Q as *u8, 8*T*D); sys_munmap(K as *u8, 8*T*D); sys_munmap(V as *u8, 8*T*D) 60 sys_munmap(ctx as *u8, 8*T*D); sys_munmap(sc as *u8, 8*T); sys_munmap(sm as *u8, 8*T) 61 return 0 62} 63 64// 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 65func vit_encoder_layer(x: *i64, wts: *i64, T: i64, D: i64, nh: i64, hd: i64, mlp: i64, eps: i64, scale: i64, out: *i64) -> i64 { 66 let ln1g: *i64 = (wts[0]) as *i64; let ln1b: *i64 = (wts[1]) as *i64 67 let Wq: *i64 = (wts[2]) as *i64; let bq: *i64 = (wts[3]) as *i64 68 let Wk: *i64 = (wts[4]) as *i64; let bk: *i64 = (wts[5]) as *i64 69 let Wv: *i64 = (wts[6]) as *i64; let bv: *i64 = (wts[7]) as *i64 70 let Wo: *i64 = (wts[8]) as *i64; let bo: *i64 = (wts[9]) as *i64 71 let ln2g: *i64 = (wts[10]) as *i64; let ln2b: *i64 = (wts[11]) as *i64 72 let fc1w: *i64 = (wts[12]) as *i64; let fc1b: *i64 = (wts[13]) as *i64 73 let fc2w: *i64 = (wts[14]) as *i64; let fc2b: *i64 = (wts[15]) as *i64 74 75 let xn: *i64 = sys_mmap(8 * T * D) as *i64 76 nx_f32_layernorm(x, ln1g, ln1b, T, D, eps, xn) 77 let attn: *i64 = sys_mmap(8 * T * D) as *i64 78 vit_mhsa(xn, Wq, bq, Wk, bk, Wv, bv, Wo, bo, T, D, nh, hd, scale, attn) 79 let hbuf: *i64 = sys_mmap(8 * T * D) as *i64 80 var i: i64 = 0 81 while i < T*D { hbuf[i] = nx_f32_add(x[i], attn[i]); i = i + 1 } 82 let hn: *i64 = sys_mmap(8 * T * D) as *i64 83 nx_f32_layernorm(hbuf, ln2g, ln2b, T, D, eps, hn) 84 let f1: *i64 = sys_mmap(8 * T * mlp) as *i64 85 nx_f32_matmul_t(hn, fc1w, f1, T, D, mlp); vit_bias_add(f1, fc1b, T, mlp) 86 nx_f32_gelu_vec(f1, T * mlp) 87 let m2: *i64 = sys_mmap(8 * T * D) as *i64 88 nx_f32_matmul_t(f1, fc2w, m2, T, mlp, D); vit_bias_add(m2, fc2b, T, D) 89 i = 0 90 while i < T*D { out[i] = nx_f32_add(hbuf[i], m2[i]); i = i + 1 } 91 sys_munmap(xn as *u8, 8*T*D); sys_munmap(attn as *u8, 8*T*D); sys_munmap(hbuf as *u8, 8*T*D) 92 sys_munmap(hn as *u8, 8*T*D); sys_munmap(f1 as *u8, 8*T*mlp); sys_munmap(m2 as *u8, 8*T*D) 93 return 0 94}