Computes the mean squared error between a policy's estimates of the expected arm payouts and the true expected payouts.
(bandit, policy)
| 24 | |
| 25 | |
| 26 | def mse(bandit, policy): |
| 27 | """ |
| 28 | Computes the mean squared error between a policy's estimates of the |
| 29 | expected arm payouts and the true expected payouts. |
| 30 | """ |
| 31 | if not hasattr(policy, "ev_estimates") or len(policy.ev_estimates) == 0: |
| 32 | return np.nan |
| 33 | |
| 34 | se = [] |
| 35 | evs = bandit.arm_evs |
| 36 | ests = sorted(policy.ev_estimates.items(), key=lambda x: x[0]) |
| 37 | for ix, (est, ev) in enumerate(zip(ests, evs)): |
| 38 | se.append((est[1] - ev) ** 2) |
| 39 | return np.mean(se) |
| 40 | |
| 41 | |
| 42 | def smooth(prev, cur, weight): |