code wiki / (root) / sketch_ucb1_test.nx

sketch_ucb1_test.nx source

↩ module page · 123 lines · 4231 B

1// sketch_ucb1_test.nx -- multi-armed bandit verification. 2 3import "syscalls.nx" 4import "sketch_ucb1.nx" 5import "sketch_types.nx" 6 7func iabs(x: i64) -> i64 { 8 if x < 0 { return -x } 9 return x 10} 11 12// Deterministic "reward" for arm i: arm 0 = 70% reward, arm 1 = 50%, arm 2 = 30%. 13// Simulated as a cyclic period of period-10: each arm produces 1.0 PPM 14// for the first (reward_rate * 10) pulls within each block. 15func reward_for(arm: i64, pull_idx: i64) -> i64 { 16 var threshold: i64 = 0 17 if arm == 0 { threshold = 7 } 18 if arm == 1 { threshold = 5 } 19 if arm == 2 { threshold = 3 } 20 let cycle: i64 = pull_idx % 10 21 if cycle < threshold { return 1000000 } 22 return 0 23} 24 25func main() -> i64 { 26 // ---- alloc ---- 27 let b: *Ucb1 = nx_ucb_alloc(3) 28 if b == (0 as *Ucb1) { return __syscall(93, 5, 0, 0, 0, 0, 0) } 29 if b.total_pulls != 0 { return __syscall(93, 6, 0, 0, 0, 0, 0) } 30 // Reject n_arms < 2. 31 if nx_ucb_alloc(1) != (0 as *Ucb1) { 32 return __syscall(93, 7, 0, 0, 0, 0, 0) 33 } 34 35 // ---- exploration: first 3 pulls cover all arms ---- 36 let pull0: i64 = nx_ucb_select(b) 37 if pull0 != 0 { return __syscall(93, 10, 0, 0, 0, 0, 0) } 38 nx_ucb_update(b, 0, 500000) // arm 0: reward 0.5 39 40 let pull1: i64 = nx_ucb_select(b) 41 if pull1 != 1 { return __syscall(93, 11, 0, 0, 0, 0, 0) } 42 nx_ucb_update(b, 1, 500000) 43 44 let pull2: i64 = nx_ucb_select(b) 45 if pull2 != 2 { return __syscall(93, 12, 0, 0, 0, 0, 0) } 46 nx_ucb_update(b, 2, 500000) 47 48 if b.total_pulls != 3 { return __syscall(93, 13, 0, 0, 0, 0, 0) } 49 50 // ---- convergence: 1000 pulls -> best arm should be detected ---- 51 let b2: *Ucb1 = nx_ucb_alloc(3) 52 var pull_count: i64 = 0 53 var i: i64 = 0 54 while i < 1000 { 55 let arm: i64 = nx_ucb_select(b2) 56 let r: i64 = reward_for(arm, pull_count) 57 nx_ucb_update(b2, arm, r) 58 pull_count = pull_count + 1 59 i = i + 1 60 } 61 if b2.total_pulls != 1000 { return __syscall(93, 20, 0, 0, 0, 0, 0) } 62 // Best arm should be 0 (70% reward). 63 if nx_ucb_best_arm(b2) != 0 { 64 return __syscall(93, 21, 0, 0, 0, 0, 0) 65 } 66 // Arm 0 should have been pulled MORE than arms 1 or 2 (exploitation). 67 let c0: i64 = nx_ucb_count(b2, 0) 68 let c1: i64 = nx_ucb_count(b2, 1) 69 let c2: i64 = nx_ucb_count(b2, 2) 70 if c0 <= c1 { return __syscall(93, 22, 0, 0, 0, 0, 0) } 71 if c0 <= c2 { return __syscall(93, 23, 0, 0, 0, 0, 0) } 72 // Mean rewards should approximate true rates. 73 // Arm 0: ~ 0.7 = 700_000 ppm. Allow +/-50_000. 74 let m0: i64 = nx_ucb_mean_ppm(b2, 0) 75 if iabs(m0 - 700000) > 50000 { 76 return __syscall(93, 24, 0, 0, 0, 0, 0) 77 } 78 let m1: i64 = nx_ucb_mean_ppm(b2, 1) 79 if iabs(m1 - 500000) > 100000 { 80 return __syscall(93, 25, 0, 0, 0, 0, 0) 81 } 82 let m2: i64 = nx_ucb_mean_ppm(b2, 2) 83 if iabs(m2 - 300000) > 100000 { 84 return __syscall(93, 26, 0, 0, 0, 0, 0) 85 } 86 87 // ---- regret bound: arm-0 fraction grows with N ---- 88 // After 1000 pulls, optimal arm pulled at least 50% of the time. 89 if c0 * 2 < 1000 { 90 return __syscall(93, 30, 0, 0, 0, 0, 0) 91 } 92 93 // ---- typed envelope ---- 94 let q: *ApproxI64 = nx_ucb_query(b2, 0) 95 if q.envelope_kind != NX_ENV_REL_STDDEV { 96 return __syscall(93, 40, 0, 0, 0, 0, 0) 97 } 98 if q.maturity != NX_MATURITY_REFERENCE_IMPL { 99 return __syscall(93, 41, 0, 0, 0, 0, 0) 100 } 101 // stderr ~ 1/sqrt(c0). 102 let isq: i64 = nx_ucb_isqrt(c0) 103 let expected_se: i64 = 1000000000 / isq 104 if iabs(q.param_a - expected_se) > 1000 { 105 return __syscall(93, 42, 0, 0, 0, 0, 0) 106 } 107 108 // ---- input validation ---- 109 if nx_ucb_update(b2, -1, 500000) != -1 { 110 return __syscall(93, 50, 0, 0, 0, 0, 0) 111 } 112 if nx_ucb_update(b2, 999, 500000) != -1 { 113 return __syscall(93, 51, 0, 0, 0, 0, 0) 114 } 115 if nx_ucb_update(b2, 0, 2000000) != -1 { // > 1.0 116 return __syscall(93, 52, 0, 0, 0, 0, 0) 117 } 118 if nx_ucb_update(b2, 0, -1) != -1 { 119 return __syscall(93, 53, 0, 0, 0, 0, 0) 120 } 121 122 return 0 123}