From 57544e5520ec03ba5c0a6790da17732b0ec3398a Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 4 Apr 2025 15:32:44 +0200 Subject: [PATCH] added ucb --- k-bandits.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/k-bandits.py b/k-bandits.py index 0c73510..f3f838d 100644 --- a/k-bandits.py +++ b/k-bandits.py @@ -20,8 +20,19 @@ def simple_max(Q, N, t, _tie_break): am = np.argmax(Q + _tie_break[t]) return am + +def ucb(Q, N, t, _tie_break): + c = 2 + if N.min() == 0: +# return np.argmax(N + _tie_break[t]) + return np.random.choice(np.flatnonzero(N == N.min())) + + M = Q + c * np.sqrt(np.divide(np.log(t), N)) + return np.argmax(M + _tie_break[t]) # breaking ties randomly +# return np.argmax(M) # breaking ties randomly + class Param: - def __init__(self, epsilon:float=0, rho:float=0, alpha:float =0, q_ic:float=0, _argmax_func=simple_max): + def __init__(self, epsilon:float=0, rho:float=0, alpha:float =0, q_ic:float=0, argmax_func=simple_max): # Anti-greediness (ability to explore) self.epsilon = epsilon @@ -36,7 +47,7 @@ class Param: self.q_ic = q_ic # argmax function - self.argmax_func=_argmax_func + self.argmax_func = argmax_func def __repr__(self): return f"epsilon={self.epsilon}, rho={self.rho}, alpha={self.alpha}, q_ic={self.q_ic}\nargmax={self.argmax_func.__repr__()}" @@ -117,7 +128,8 @@ if __name__ == '__main__': pl.figure(figsize=(12, 8)) # params = [Param(epsilon=0.0), Param(epsilon=0.01), Param(epsilon=0.1)] - params = [Param(epsilon=0.1), Param(alpha=0.2, q_ic=5)] +# params = [Param(epsilon=0.1), Param(alpha=0.2, q_ic=5)] + params = [Param(epsilon=0.1), Param(argmax_func=ucb)] legend = [] for param in params: res_r, res_a = batch(k_arms, episode_len, num_realisations, param)