File size: 4,682 Bytes
f916dc0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | #!/usr/bin/env python3
"""
example.py — Usage demonstrations for hv_drift_rl.
NumPy only.
"""
import numpy as np
from hv_drift_rl import (
Corridor, make_obs_drift, make_rew_drift,
make_mechanism, train, ALL_MECHANISMS, ALL_SHAPES,
)
def demo_single_run():
print("=" * 72)
print("Demo 1 — single run, one mechanism, one shape")
print("=" * 72)
print()
rng = np.random.default_rng(42)
obs_drift = make_obs_drift('linear', rng)
rew_drift = make_rew_drift('linear', rng)
env = Corridor(length=20, obs_drift_fn=obs_drift, rew_drift_fn=rew_drift)
mech, disc = make_mechanism('delta_obs_rew')
res = train(env, mech, disc, n_episodes=500, seed=42)
print(f" mechanism: delta_obs_rew")
print(f" shape: linear")
print(f" success: {res['success_rate']:.1%}")
print(f" mean rew: {np.mean(res['rewards'][-100:]):.3f}")
print(f" mean steps: {np.mean(res['steps'][-100:]):.1f}")
print()
def demo_ablation():
print("=" * 72)
print("Demo 2 — the ablation: delta_obs vs delta_rew_only vs delta_obs_rew")
print("=" * 72)
print()
print(f" {'shape':<16s} {'delta_obs':>12s} {'delta_rew_only':>15s} "
f"{'delta_obs_rew':>15s}")
print(" " + "-" * 62)
for shape in ['none', 'offset', 'linear']:
row = []
for mech_name in ['delta_obs', 'delta_rew_only', 'delta_obs_rew']:
rng = np.random.default_rng(42)
obs_drift = make_obs_drift(shape, rng)
rew_drift = make_rew_drift(shape, rng)
env = Corridor(length=20, obs_drift_fn=obs_drift,
rew_drift_fn=rew_drift)
mech, disc = make_mechanism(mech_name)
res = train(env, mech, disc, n_episodes=500, seed=42)
row.append(res['success_rate'])
print(f" {shape:<16s} {row[0]:>12.1%} {row[1]:>15.1%} "
f"{row[2]:>15.1%}")
print()
print(" Only the composite rescues the failing shapes.")
print()
def demo_all_mechanisms():
print("=" * 72)
print("Demo 3 — all mechanisms on one shape (linear)")
print("=" * 72)
print()
print(f" {'mechanism':<20s} {'success':>10s} {'preserves_terminal':>20s}")
print(" " + "-" * 54)
terminal_ok = {
'no_correction': True, 'delta_obs': True, 'delta_rew_only': True,
'delta_obs_rew': True, 'safe_baseline': True, 'gym_rms': True,
'popart': False, 'reward_clipping': False, 'anchoring': True,
}
for mech_name in ALL_MECHANISMS:
rng = np.random.default_rng(42)
obs_drift = make_obs_drift('linear', rng)
rew_drift = make_rew_drift('linear', rng)
env = Corridor(length=20, obs_drift_fn=obs_drift,
rew_drift_fn=rew_drift)
mech, disc = make_mechanism(mech_name)
res = train(env, mech, disc, n_episodes=500, seed=42)
print(f" {mech_name:<20s} {res['success_rate']:>10.1%} "
f"{str(terminal_ok[mech_name]):>20s}")
print()
print(" Note: the two failing mechanisms (popart, reward_clipping)")
print(" are also the two that don't preserve terminal reward.")
print()
def demo_drift_shapes():
print("=" * 72)
print("Demo 4 — best mechanism (delta_obs_rew) across all shapes")
print("=" * 72)
print()
print(f" {'shape':<20s} {'success':>10s}")
print(" " + "-" * 34)
for shape in ALL_SHAPES:
rng = np.random.default_rng(42)
obs_drift = make_obs_drift(shape, rng)
rew_drift = make_rew_drift(shape, rng)
env = Corridor(length=20, obs_drift_fn=obs_drift,
rew_drift_fn=rew_drift)
mech, disc = make_mechanism('delta_obs_rew')
res = train(env, mech, disc, n_episodes=500, seed=42)
print(f" {shape:<20s} {res['success_rate']:>10.1%}")
print()
def demo_self_drift():
print("=" * 72)
print("Demo 5 — self-drift: the non-monotonic decay curve")
print("=" * 72)
print()
from hv_drift_rl import train_self
print(f" {'rate':>7s} {'blur':>10s} {'decay':>10s}")
print(" " + "-" * 30)
for rate in [0.0, 0.001, 0.005, 0.01, 0.02, 0.05]:
s_blur = train_self('blur', rate)
s_decay = train_self('decay', rate)
print(f" {rate:>7.3f} {s_blur:>10.1%} {s_decay:>10.1%}")
print()
print(" Blur never fails. Decay is worst at 0.01, then partially")
print(" recovers at 0.02-0.05 (homogenization toward Q.mean).")
print()
def main():
demo_single_run()
demo_ablation()
demo_all_mechanisms()
demo_drift_shapes()
demo_self_drift()
if __name__ == '__main__':
main() |