nx_olmoe_recon_gate.nx source
↩ module page · 65 lines · 2512 B
1// nx_olmoe_recon_gate.nx -- recon: enumerate blk.0.* tensor names + dims + types in OLMoE so the rung-4
2// attention (QK-norm) is wired to the model's ACTUAL shape, not assumption. Prints every blk.0 tensor.
3// expect_exit: 0 license_tier: ORIGINAL
4import "nx_syscalls.nx"
5import "nx_tier.nx"
6import "nx_le.nx"
7import "nx_tensor.nx"
8import "nx_gguf.nx"
9
10func rc_w(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 }
11func rc_n(v: i64) -> i64 {
12 var m: i64 = v
13 if m < 0 { rc_w("-" as *u8); m = 0 - m }
14 let t: *u8 = sys_mmap(24)
15 var k: i64 = 0
16 if m == 0 { t[0] = 48 as u8; k = 1 }
17 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
18 let o: *u8 = sys_mmap(24)
19 var i: i64 = 0
20 while i < k { o[i] = t[k - 1 - i]; i = i + 1 }
21 sys_write(1, o, k)
22 return 0
23}
24// true if ti.name starts with the 6 bytes "blk.0."
25func rc_is_blk0(ti: *NxGgufTensorInfo) -> i64 {
26 if ti.name_len < 6 { return 0 }
27 let nm: *u8 = ti.name
28 if nm[0]!=(98 as u8) { return 0 } // b
29 if nm[1]!=(108 as u8) { return 0 } // l
30 if nm[2]!=(107 as u8) { return 0 } // k
31 if nm[3]!=(46 as u8) { return 0 } // .
32 if nm[4]!=(48 as u8) { return 0 } // 0
33 if nm[5]!=(46 as u8) { return 0 } // .
34 return 1
35}
36
37func main() -> i64 {
38 rc_w("=== OLMoE blk.0 tensor recon ===\n" as *u8)
39 let ln: *i64 = sys_mmap(8) as *i64
40 ln[0] = 0
41 let buf: *u8 = sys_map_file("/home/elderwesto/nx_stage/nx_moe_model.gguf" as *u8, ln)
42 if (buf as i64) == 0 { rc_w("MODEL ABSENT\n" as *u8); return 1 }
43 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader
44 if nx_gguf_parse(buf, ln[0], hdr) != NX_GGUF_OK { rc_w("PARSE FAIL\n" as *u8); return 1 }
45 let hloc: *NxGgufHeader = hdr
46 var i: i64 = 0
47 var shown: i64 = 0
48 while i < hloc.tensor_count {
49 let ti: *NxGgufTensorInfo = nx_gguf_tensor_at(hloc, i)
50 if rc_is_blk0(ti) == 1 {
51 rc_w(" " as *u8)
52 sys_write(1, ti.name, ti.name_len)
53 rc_w(" dims=" as *u8); rc_n(ti.dim_0)
54 if ti.n_dims >= 2 { rc_w("x" as *u8); rc_n(ti.dim_1) }
55 if ti.n_dims >= 3 { rc_w("x" as *u8); rc_n(ti.dim_2) }
56 rc_w(" ndims=" as *u8); rc_n(ti.n_dims)
57 rc_w(" ty=" as *u8); rc_n(ti.ggml_type)
58 rc_w("\n" as *u8)
59 shown = shown + 1
60 }
61 i = i + 1
62 }
63 rc_w("blk.0 tensors=" as *u8); rc_n(shown); rc_w("\n" as *u8)
64 return 0
65}