nx_f32_kv_cache_test.nx source
↩ module page · 118 lines · 4215 B
1// nx_f32_kv_cache_test.nx -- smoke for nx_f32_kv_cache.nx.
2
3import "nx_syscalls.nx"
4import "nx_tier.nx"
5import "nx_f32_kv_cache.nx"
6
7func main() -> i64 {
8 var vi: nx_int = 0
9 while vi < NX_F32_KVC_N_VERDICTS {
10 if nx_f32_kvc_verdict_is_valid(vi) != 1 { return 5 + vi }
11 vi = vi + 1
12 }
13
14 // ===== Allocate cache: n_layers=2, n_kv_heads=2, max_seq=4, head_dim=2 =====
15 // kv_dim = 4
16 let c: *NxF32KVCache = nx_f32_kv_cache_alloc(2, 2, 4, 2)
17 if c == (0 as *NxF32KVCache) { return 10 }
18 if c.n_layers != 2 { return 11 }
19 if c.n_kv_heads != 2 { return 12 }
20 if c.max_seq_len != 4 { return 13 }
21 if c.head_dim != 2 { return 14 }
22 if nx_f32_kv_cache_get_seq_len(c) != 0 { return 15 }
23
24 // ===== Append 1 token for layer 0 =====
25 let K_new1: *i64 = sys_mmap(4 * 8) as *i64
26 let V_new1: *i64 = sys_mmap(4 * 8) as *i64
27 K_new1[0] = 0x3F800000 // 1.0
28 K_new1[1] = 0x40000000 // 2.0
29 K_new1[2] = 0x40400000 // 3.0
30 K_new1[3] = 0x40800000 // 4.0
31 V_new1[0] = 0x40A00000 // 5.0
32 V_new1[1] = 0x40C00000 // 6.0
33 V_new1[2] = 0x40E00000 // 7.0
34 V_new1[3] = 0x41000000 // 8.0
35
36 let v_app: nx_int = nx_f32_kv_cache_append_layer(c, 0, K_new1, V_new1, 1)
37 if v_app != NX_F32_KVC_OK { return 20 + v_app }
38
39 // Append same K/V for layer 1.
40 let K_new1b: *i64 = sys_mmap(4 * 8) as *i64
41 let V_new1b: *i64 = sys_mmap(4 * 8) as *i64
42 K_new1b[0] = 0x41200000 // 10.0
43 K_new1b[1] = 0x41200000
44 K_new1b[2] = 0x41200000
45 K_new1b[3] = 0x41200000
46 V_new1b[0] = 0x41400000 // 12.0
47 V_new1b[1] = 0x41400000
48 V_new1b[2] = 0x41400000
49 V_new1b[3] = 0x41400000
50
51 let v_app2: nx_int = nx_f32_kv_cache_append_layer(c, 1, K_new1b, V_new1b, 1)
52 if v_app2 != NX_F32_KVC_OK { return 30 }
53
54 // Advance seq_len after BOTH layers appended.
55 let v_adv: nx_int = nx_f32_kv_cache_advance(c, 1)
56 if v_adv != NX_F32_KVC_OK { return 31 }
57 if nx_f32_kv_cache_get_seq_len(c) != 1 { return 32 }
58
59 // ===== Verify layer 0 K/V contents =====
60 let K0: *i64 = nx_f32_kv_cache_get_K_layer(c, 0)
61 if K0[0] != 0x3F800000 { return 40 }
62 if K0[1] != 0x40000000 { return 41 }
63 if K0[2] != 0x40400000 { return 42 }
64 if K0[3] != 0x40800000 { return 43 }
65
66 let V0: *i64 = nx_f32_kv_cache_get_V_layer(c, 0)
67 if V0[0] != 0x40A00000 { return 50 }
68 if V0[1] != 0x40C00000 { return 51 }
69 if V0[2] != 0x40E00000 { return 52 }
70 if V0[3] != 0x41000000 { return 53 }
71
72 // Verify layer 1 isolation
73 let K1: *i64 = nx_f32_kv_cache_get_K_layer(c, 1)
74 if K1[0] != 0x41200000 { return 60 }
75 if K1[1] != 0x41200000 { return 61 }
76
77 // ===== Append a second row for layer 0 =====
78 let K_new2: *i64 = sys_mmap(4 * 8) as *i64
79 let V_new2: *i64 = sys_mmap(4 * 8) as *i64
80 K_new2[0] = 0x41700000 // 15.0
81 K_new2[1] = 0x41700000
82 K_new2[2] = 0x41700000
83 K_new2[3] = 0x41700000
84 V_new2[0] = 0x41800000 // 16.0
85 V_new2[1] = 0x41800000
86 V_new2[2] = 0x41800000
87 V_new2[3] = 0x41800000
88
89 let v_app3: nx_int = nx_f32_kv_cache_append_layer(c, 0, K_new2, V_new2, 1)
90 if v_app3 != NX_F32_KVC_OK { return 70 }
91 nx_f32_kv_cache_advance(c, 1)
92 if nx_f32_kv_cache_get_seq_len(c) != 2 { return 71 }
93
94 // Verify both rows exist in layer 0 K
95 let K0_after: *i64 = nx_f32_kv_cache_get_K_layer(c, 0)
96 // Row 0 (originally appended)
97 if K0_after[0] != 0x3F800000 { return 80 }
98 if K0_after[3] != 0x40800000 { return 81 }
99 // Row 1 (newly appended)
100 if K0_after[4] != 0x41700000 { return 82 }
101 if K0_after[7] != 0x41700000 { return 83 }
102
103 // ===== Bad-layer rejection =====
104 let v_bad: nx_int = nx_f32_kv_cache_append_layer(c, 5, K_new1, V_new1, 1)
105 if v_bad != NX_F32_KVC_ERR_BAD_LAYER { return 90 }
106
107 // ===== Overflow: append until max_seq_len + 1 =====
108 // We already have seq_len=2, max=4. Try to advance by 3 -> overflow.
109 let v_ovf: nx_int = nx_f32_kv_cache_advance(c, 3)
110 if v_ovf != NX_F32_KVC_ERR_OVERFLOW { return 100 }
111
112 // ===== Reset =====
113 let v_rst: nx_int = nx_f32_kv_cache_reset(c)
114 if v_rst != NX_F32_KVC_OK { return 110 }
115 if nx_f32_kv_cache_get_seq_len(c) != 0 { return 111 }
116
117 return 0
118}