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}