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}