hv-drift-rl / example.py
zeechimp's picture
Create example.py
f916dc0 verified
Raw History Blame Contribute Delete
4.68 kB
#!/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()