code wiki / (root) / nx_f32_kv_cache_test.nx

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}