"""Learning from your hand: record demonstrations, then clone them.

Reinforcement learning from scratch has to stumble onto a good action by
accident before it can be rewarded for it. Balancing is unforgiving about
that - almost every random policy falls in half a second, so nearly every
episode looks equally bad and there is very little signal to climb.

Showing it what to do removes that problem. You drive the wheels by hand, we
record what the robot could see and what you did, and the policy is trained by
plain supervised regression to copy you. That is behaviour cloning. It will not
be as good as you, and it will fail in states you never visited, but it starts
PPO somewhere sensible instead of at noise.

  python demo.py --list
  python demo.py --clone johnny6_balance --task balance --out ../runs/seed
"""

import argparse
import json
import time
from pathlib import Path

import numpy as np

from env import RobotEnv
from ppo import PPO, Policy

ROOT = Path(__file__).parent.parent
DEMOS = ROOT / "demos"


def save(name, task, obs, act, meta=None):
    DEMOS.mkdir(parents=True, exist_ok=True)
    p = DEMOS / f"{name}.json"
    old = {"obs": [], "act": []}
    if p.exists():
        try:
            old = json.load(open(p))
        except (OSError, ValueError):
            pass
    data = {
        "task": task,
        "obs": old.get("obs", []) + [list(map(float, o)) for o in obs],
        "act": old.get("act", []) + [list(map(float, a)) for a in act],
        "meta": meta or {},
        "updated": time.time(),
    }
    json.dump(data, open(p, "w"))
    return p, len(data["obs"])


def load(name):
    p = DEMOS / f"{name}.json"
    d = json.load(open(p))
    return np.array(d["obs"], dtype=float), np.array(d["act"], dtype=float), d["task"]


def clone(name, task, spec, out, epochs=300, lr=2e-3, hidden=(64, 64), seed=0):
    """Fit a policy to the recorded demonstrations by supervised regression."""
    obs, act, demo_task = load(name)
    task = task or demo_task
    env = RobotEnv(spec, seed=seed, task=task)
    if obs.shape[1] != env.obs_dim or act.shape[1] != env.act_dim:
        raise SystemExit(
            f"demonstration shape {obs.shape[1]}x{act.shape[1]} does not match "
            f"the {task!r} task ({env.obs_dim}x{env.act_dim}). Re-record it.")

    policy = Policy(env.obs_dim, env.act_dim, hidden=hidden, seed=seed)
    policy.norm.update(obs)
    x = policy.norm(obs)

    from nets import Adam
    opt = Adam(policy.pi.params(), lr=lr)
    n = len(x)
    rng = np.random.default_rng(seed)
    idx = np.arange(n)
    mb = max(32, n // 20)
    print(f"cloning {n} demonstrated steps for {task!r}")
    for ep in range(epochs):
        rng.shuffle(idx)
        tot = 0.0
        for s in range(0, n, mb):
            b = idx[s:s + mb]
            cache = []
            mu = policy.pi.forward(x[b], cache)
            err = mu - act[b]
            gW, gb = policy.pi.backward(cache, 2.0 * err / len(b))
            opt.step(gW + gb, max_norm=2.0)
            tot += float(np.mean(err ** 2)) * len(b)
        if ep % 50 == 0 or ep == epochs - 1:
            print(f"  epoch {ep:4d}  mean squared error {tot / n:.5f}")

    # start PPO with a wide-ish spread so it still explores around your example
    policy.log_std[:] = -1.0
    out = Path(out)
    out.mkdir(parents=True, exist_ok=True)
    policy.save(out / "policy_best.json",
                meta={"spec": str(Path(spec).resolve()), "task": task,
                      "steps": 0, "obs_layout": "union",
                      "cloned_from": name, "demo_steps": n})
    policy.save(out / "policy_latest.json",
                meta={"spec": str(Path(spec).resolve()), "task": task,
                      "steps": 0, "obs_layout": "union",
                      "cloned_from": name, "demo_steps": n})
    print(f"\nwrote {out}/policy_best.json")
    return policy


def evaluate(policy, spec, task, episodes=10, seed=9000):
    ups = []
    for ep in range(episodes):
        env = RobotEnv(spec, seed=seed + ep, task=task)
        rng = np.random.default_rng(seed + ep)
        obs = env.reset()
        for _ in range(env.max_steps):
            a, _, _ = policy.act(obs[None], rng, deterministic=True)
            obs, r, fell, trunc, info = env.step(np.clip(a[0], -1, 1))
            if fell:
                break
        ups.append(env.t)
    return np.array(ups)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--list", action="store_true")
    ap.add_argument("--clone", default=None, help="demonstration name")
    ap.add_argument("--task", default=None)
    ap.add_argument("--spec", default=str(ROOT / "robots/wheeled_biped.json"))
    ap.add_argument("--out", default=str(ROOT / "runs/johnny6__balance"))
    ap.add_argument("--epochs", type=int, default=300)
    args = ap.parse_args()

    if args.list:
        DEMOS.mkdir(parents=True, exist_ok=True)
        for p in sorted(DEMOS.glob("*.json")):
            d = json.load(open(p))
            print(f"  {p.stem:24s} {len(d['obs']):6d} steps   task {d['task']}")
        return

    if args.clone:
        pol = clone(args.clone, args.task, args.spec, args.out, epochs=args.epochs)
        env = RobotEnv(args.spec, seed=0, task=args.task or load(args.clone)[2])
        ups = evaluate(pol, args.spec, env.task_name)
        full = env.max_steps * env.control_dt
        print(f"cloned policy: upright {ups.mean():.2f}s of {full:.0f}s, "
              f"{int((ups >= full - 1e-6).sum())}/{len(ups)} full episodes")


if __name__ == "__main__":
    main()
