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}