code wiki / (root) / sketch_epsilon_greedy_test.nx

sketch_epsilon_greedy_test.nx source

↩ module page · 152 lines · 5084 B

1// sketch_epsilon_greedy_test.nx -- epsilon-greedy bandit verification. 2 3import "syscalls.nx" 4import "sketch_epsilon_greedy.nx" 5import "sketch_types.nx" 6 7func iabs(x: i64) -> i64 { 8 if x < 0 { return -x } 9 return x 10} 11 12func reward_for(arm: i64, pull_idx: i64) -> i64 { 13 var threshold: i64 = 0 14 if arm == 0 { threshold = 7 } 15 if arm == 1 { threshold = 5 } 16 if arm == 2 { threshold = 3 } 17 let cycle: i64 = pull_idx % 10 18 if cycle < threshold { return 1000000 } 19 return 0 20} 21 22func main() -> i64 { 23 // ---- alloc + parameter validation ---- 24 let eg: *EpsilonGreedy = nx_eg_alloc(3, NX_EG_MODE_FIXED, 100000, 42) 25 if eg == (0 as *EpsilonGreedy) { return __syscall(93, 5, 0, 0, 0, 0, 0) } 26 // Reject mode out of range. 27 if nx_eg_alloc(3, 99, 100000, 1) != (0 as *EpsilonGreedy) { 28 return __syscall(93, 6, 0, 0, 0, 0, 0) 29 } 30 // Reject n_arms < 2. 31 if nx_eg_alloc(1, NX_EG_MODE_FIXED, 100000, 1) != (0 as *EpsilonGreedy) { 32 return __syscall(93, 7, 0, 0, 0, 0, 0) 33 } 34 // Reject epsilon > 1.0. 35 if nx_eg_alloc(3, NX_EG_MODE_FIXED, 2000000, 1) != (0 as *EpsilonGreedy) { 36 return __syscall(93, 8, 0, 0, 0, 0, 0) 37 } 38 39 // ---- convergence (fixed epsilon=0.1, 1000 pulls, true rates 70/50/30) ---- 40 var pull_idx: i64 = 0 41 var i: i64 = 0 42 while i < 1000 { 43 let arm: i64 = nx_eg_select(eg) 44 let r: i64 = reward_for(arm, pull_idx) 45 nx_eg_update(eg, arm, r) 46 pull_idx = pull_idx + 1 47 i = i + 1 48 } 49 if eg.total_pulls != 1000 { return __syscall(93, 10, 0, 0, 0, 0, 0) } 50 // Best arm should be 0. 51 if nx_eg_best_arm(eg) != 0 { 52 return __syscall(93, 11, 0, 0, 0, 0, 0) 53 } 54 // Mean reward of arm 0 should be near 0.7. 55 let m0: i64 = nx_eg_mean_ppm(eg, 0) 56 if iabs(m0 - 700000) > 60000 { 57 return __syscall(93, 12, 0, 0, 0, 0, 0) 58 } 59 // Arm 0 should be pulled more than arms 1 and 2 (exploitation). 60 let c0: i64 = nx_eg_count(eg, 0) 61 let c1: i64 = nx_eg_count(eg, 1) 62 let c2: i64 = nx_eg_count(eg, 2) 63 if c0 <= c1 { return __syscall(93, 13, 0, 0, 0, 0, 0) } 64 if c0 <= c2 { return __syscall(93, 14, 0, 0, 0, 0, 0) } 65 66 // ---- determinism ---- 67 let eg_a: *EpsilonGreedy = nx_eg_alloc(3, NX_EG_MODE_FIXED, 100000, 99) 68 let eg_b: *EpsilonGreedy = nx_eg_alloc(3, NX_EG_MODE_FIXED, 100000, 99) 69 pull_idx = 0 70 i = 0 71 while i < 200 { 72 let aa: i64 = nx_eg_select(eg_a) 73 let ab: i64 = nx_eg_select(eg_b) 74 if aa != ab { 75 return __syscall(93, 20, 0, 0, 0, 0, 0) 76 } 77 let ra: i64 = reward_for(aa, pull_idx) 78 nx_eg_update(eg_a, aa, ra) 79 nx_eg_update(eg_b, ab, ra) 80 pull_idx = pull_idx + 1 81 i = i + 1 82 } 83 // Same total pulls and same counts. 84 if eg_a.total_pulls != eg_b.total_pulls { 85 return __syscall(93, 21, 0, 0, 0, 0, 0) 86 } 87 i = 0 88 while i < 3 { 89 if nx_eg_count(eg_a, i) != nx_eg_count(eg_b, i) { 90 return __syscall(93, 22, 0, 0, 0, 0, 0) 91 } 92 i = i + 1 93 } 94 95 // ---- decay mode ---- 96 let eg_d: *EpsilonGreedy = nx_eg_alloc(3, NX_EG_MODE_DECAY_LINEAR, 1000000, 7) 97 // At t=0, epsilon = 1.0 / 1 = 1.0 (100% explore). 98 let e0: i64 = nx_eg_current_epsilon_ppm(eg_d) 99 if e0 != 1000000 { 100 return __syscall(93, 30, 0, 0, 0, 0, 0) 101 } 102 // Force-update 100 times. 103 pull_idx = 0 104 i = 0 105 while i < 100 { 106 let arm: i64 = nx_eg_select(eg_d) 107 let r: i64 = reward_for(arm, pull_idx) 108 nx_eg_update(eg_d, arm, r) 109 pull_idx = pull_idx + 1 110 i = i + 1 111 } 112 // Now epsilon ~ 1.0/101 ~ 9900. 113 let e100: i64 = nx_eg_current_epsilon_ppm(eg_d) 114 if iabs(e100 - 9900) > 500 { 115 return __syscall(93, 31, 0, 0, 0, 0, 0) 116 } 117 118 // ---- epsilon = 0 means pure exploit ---- 119 let eg_pure: *EpsilonGreedy = nx_eg_alloc(3, NX_EG_MODE_FIXED, 0, 7) 120 // First select on empty bandit: best_arm_by_mean returns random (no data). 121 let first: i64 = nx_eg_select(eg_pure) 122 nx_eg_update(eg_pure, first, 1000000) // perfect reward 123 // From now on, eg_pure should always pull `first` (highest mean = 1.0). 124 i = 0 125 while i < 50 { 126 let arm: i64 = nx_eg_select(eg_pure) 127 if arm != first { 128 return __syscall(93, 40, 0, 0, 0, 0, 0) 129 } 130 nx_eg_update(eg_pure, arm, 1000000) 131 i = i + 1 132 } 133 134 // ---- typed envelope ---- 135 let q: *ApproxI64 = nx_eg_query(eg, 0) 136 if q.envelope_kind != NX_ENV_REL_STDDEV { 137 return __syscall(93, 50, 0, 0, 0, 0, 0) 138 } 139 if q.maturity != NX_MATURITY_REFERENCE_IMPL { 140 return __syscall(93, 51, 0, 0, 0, 0, 0) 141 } 142 143 // ---- input validation ---- 144 if nx_eg_update(eg, -1, 500000) != -1 { 145 return __syscall(93, 60, 0, 0, 0, 0, 0) 146 } 147 if nx_eg_update(eg, 0, 2000000) != -1 { 148 return __syscall(93, 61, 0, 0, 0, 0, 0) 149 } 150 151 return 0 152}