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}