"""Addestramento di Pulcino con Augmented Random Search (ARS V2-t), puro numpy + multiprocessing.

ARS (Mania, Guy, Recht 2018) in breve, per ogni iterazione:
  1. estrai N direzioni casuali δ_k (rumore gaussiano della stessa forma dei parametri θ)
  2. prova θ + σ·δ_k e θ − σ·δ_k in simulazione ("antitetiche"), ottieni ricompense r+_k, r−_k
  3. tieni solo le b direzioni migliori (max(r+, r−) più alto)        ← la "t" di V2-t
  4. θ ← θ + α / (b·σ_R) · Σ (r+_k − r−_k) δ_k    (σ_R = dev. std. delle 2b ricompense usate)
  5. aggiorna media/dev. std. delle osservazioni viste (normalizzazione online) ← la "V2"
Niente gradienti, niente GPU, niente backprop: solo tante simulazioni in parallelo sui core
della CPU. Perfetto per un Mac o un portatile.

Politica: MLP 16 → 32 → 32 → 6 con tanh (esattamente come policy.h in SPEC §5.2).
Scelta documentata: alleniamo DIRETTAMENTE l'MLP (non prima una lineare). Lo strato di uscita
parte da zero, quindi all'inizio la politica comanda la posa "stand" (azione 0) e ARS la
"sblocca" poco a poco. Una politica lineare non è rappresentabile esattamente dal formato
policy.h (che ha due strati nascosti tanh), quindi l'MLP diretto evita conversioni.

Modalità --cpg: stesso algoritmo, ma i parametri sono i 12 numeri di gait.json (SPEC §5.1).

Esempi:
    python train_ars.py --iters 350 --workers 6               # politica MLP (~3-4 min)
    python train_ars.py --cpg --iters 150 --workers 6         # gait.json con ES (~1 min)
    python train_ars.py --resume out/ars_mlp.npz --iters 100  # riprende da checkpoint
"""
import argparse
import json
import multiprocessing as mp
import os
import time
import tempfile
from contextlib import nullcontext

import numpy as np

import cpg as cpgmod
from env import N_ACT, N_OBS, PulcinoEnv

QUI = os.path.dirname(os.path.abspath(__file__))
OUT = os.path.join(QUI, "out")
HID = 32
STD_MIN = 0.05       # dev. std. minima per la normalizzazione (evita divisioni per ~0)
CLIP_OBS = 5.0
RAMPA_CPG = 1.0     # s, rampa del comando all'avvio del CPG


# =========================== la rete (identica a policy.h) ===========================
FORME = [("W1", (HID, N_OBS)), ("B1", (HID,)), ("W2", (HID, HID)), ("B2", (HID,)),
         ("W3", (N_ACT, HID)), ("B3", (N_ACT,))]
N_PARAM_MLP = sum(int(np.prod(s)) for _, s in FORME)


def spacchetta(theta):
    """Vettore piatto θ → dict di matrici W[out][in] e bias (row-major, come in C)."""
    out, i = {}, 0
    for nome, forma in FORME:
        n = int(np.prod(forma))
        out[nome] = theta[i:i + n].reshape(forma)
        i += n
    return out


def theta_iniziale(rng):
    """Strati nascosti con init "alla Xavier", uscita a zero (azione iniziale = posa stand)."""
    p = []
    for nome, forma in FORME:
        if nome in ("W1", "W2"):
            p.append(rng.normal(0, 1.0 / np.sqrt(forma[1]), size=forma).ravel())
        else:
            p.append(np.zeros(int(np.prod(forma))))
    return np.concatenate(p)


def mlp(theta, obs, mean, std):
    """Politica: normalizza, clip ±5, tre strati tanh. Restituisce a ∈ [-1,1]^6."""
    P = spacchetta(theta)
    x = np.clip((obs - mean) / std, -CLIP_OBS, CLIP_OBS)
    h = np.tanh(P["W1"] @ x + P["B1"])
    h = np.tanh(P["W2"] @ h + P["B2"])
    return np.tanh(P["W3"] @ h + P["B3"])


# =========================== rollout (nei processi figli) ===========================
_ENV = {}


def _env(tipo, episode_s, cmd_mode):
    chiave = (tipo, episode_s, cmd_mode)
    if chiave not in _ENV:
        _ENV[chiave] = PulcinoEnv(seed=0, episode_s=episode_s, cmd_mode=cmd_mode,
                                  peso_clock=0.0 if tipo == "cpg" else 0.3)
    return _ENV[chiave]


def rollout(tipo, theta, mean, std, env_seed, episode_s, cmd_mode, raccogli_obs=False):
    """Un episodio. Restituisce (ricompensa totale, passi, caduto, statistiche obs)."""
    env = _env(tipo, episode_s, cmd_mode)
    obs = env.reset(seed=env_seed, cmd=[1.0, 0.0] if tipo == "cpg" and cmd_mode == "avanti" else None)
    tot, n = 0.0, 0
    s1 = np.zeros(N_OBS); s2 = np.zeros(N_OBS)
    gait = cpgmod.CPG(cpgmod.vettore_a_gait(theta)) if tipo == "cpg" else None
    done, info = False, {"caduto": False}
    while not done:
        if raccogli_obs:
            s1 += obs; s2 += obs * obs
        if tipo == "cpg":
            roll, pitch = env.imu_roll_pitch(env.g_last)
            # rampa di 1 s sul comando: il CPG non parte "a scalino" (consigliata anche nel firmware)
            rampa = min(1.0, env.t / RAMPA_CPG)
            q = gait.q(env.t, vx=env.cmd[0] * rampa, wz=env.cmd[1] * rampa,
                       roll_imu=roll, pitch_imu=pitch)
            obs, r, done, info = env.step_q(q)
        else:
            obs, r, done, info = env.step(mlp(theta, obs, mean, std))
        tot += r; n += 1
    return tot, n, info["caduto"], (s1, s2, n if raccogli_obs else 0)


def _lavoro(args):
    """Valuta una coppia antitetica θ ± σδ con LO STESSO seed d'ambiente (numeri casuali comuni:
    riduce molto la varianza della stima)."""
    tipo, theta, mean, std, dseed, sigma, env_seed, episode_s, cmd_mode = args
    delta = np.random.default_rng(dseed).standard_normal(theta.size)
    rp = rollout(tipo, theta + sigma * delta, mean, std, env_seed, episode_s, cmd_mode, True)
    rm = rollout(tipo, theta - sigma * delta, mean, std, env_seed, episode_s, cmd_mode, True)
    return rp, rm


def _valuta(args):
    tipo, theta, mean, std, env_seed, episode_s, cmd_mode = args
    r, n, caduto, _ = rollout(tipo, theta, mean, std, env_seed, episode_s, cmd_mode)
    return r, n, caduto


# =========================== normalizzazione online ===========================
class Normalizzatore:
    def __init__(self):
        self.n = 0
        self.s1 = np.zeros(N_OBS)
        self.s2 = np.zeros(N_OBS)

    def aggiungi(self, s1, s2, n):
        self.s1 += s1; self.s2 += s2; self.n += n

    def mean_std(self):
        if self.n < 2:
            return np.zeros(N_OBS), np.ones(N_OBS)
        m = self.s1 / self.n
        v = np.maximum(self.s2 / self.n - m * m, 0.0)
        return m, np.maximum(np.sqrt(v), STD_MIN)


# =========================== ciclo principale ===========================
def salva(path, tipo, theta, norm, it, args, storico, rng=None, migliore=-np.inf):
    """Checkpoint atomico: un'interruzione non tronca l'ultimo file valido."""
    mean, std = norm.mean_std()
    destination = os.path.abspath(path)
    os.makedirs(os.path.dirname(destination), exist_ok=True)
    fd, temporary = tempfile.mkstemp(prefix='.checkpoint-', suffix='.npz', dir=os.path.dirname(destination))
    try:
        with os.fdopen(fd, 'wb') as stream:
            np.savez(stream, tipo=tipo, theta=theta, obs_mean=mean, obs_std=std,
                     norm_n=norm.n, norm_s1=norm.s1, norm_s2=norm.s2, iterazione=it,
                     args=json.dumps(vars(args)), storico=np.array(storico),
                     rng_state=json.dumps(rng.bit_generator.state) if rng is not None else '',
                     migliore=migliore)
            stream.flush(); os.fsync(stream.fileno())
        os.replace(temporary, destination)
    finally:
        if os.path.exists(temporary): os.unlink(temporary)


def main():
    ap = argparse.ArgumentParser(description="ARS V2-t per Pulcino (MLP o CPG)")
    ap.add_argument("--cpg", action="store_true", help="ottimizza gait.json invece dell'MLP")
    ap.add_argument("--iters", type=int, default=300)
    ap.add_argument("--dirs", type=int, default=32, help="N direzioni per iterazione")
    ap.add_argument("--top", type=int, default=16, help="b direzioni migliori usate")
    ap.add_argument("--sigma", type=float, default=None, help="ampiezza del rumore di esplorazione")
    ap.add_argument("--lr", type=float, default=None, help="passo di aggiornamento α")
    ap.add_argument("--episode", type=float, default=6.0, help="durata episodio (s)")
    ap.add_argument("--cmd", default="avanti", choices=["avanti", "tutti", "fermo"])
    ap.add_argument("--workers", type=int, default=max(1, min(8, (os.cpu_count() or 2) - 2)))
    ap.add_argument("--eval-every", type=int, default=10)
    ap.add_argument("--eval-eps", type=int, default=6)
    ap.add_argument("--seed", type=int, default=1)
    ap.add_argument("--resume", default=None, help="checkpoint .npz da cui ripartire")
    ap.add_argument("--out", default=None, help="checkpoint di uscita (default out/ars_<tipo>.npz)")
    ap.add_argument('--checkpoint-every', type=int, default=10, help='salvataggio ogni N iterazioni')
    args = ap.parse_args()
    if not 1 <= args.top <= args.dirs or min(args.iters, args.workers, args.eval_every, args.eval_eps, args.checkpoint_every) < 1:
        ap.error('itera/processi/intervalli positivi; 1 <= top <= dirs')

    tipo = "cpg" if args.cpg else "mlp"
    sigma = args.sigma if args.sigma is not None else (0.04 if tipo == "cpg" else 0.02)
    lr = args.lr if args.lr is not None else 0.02
    os.makedirs(OUT, exist_ok=True)
    out_path = args.out or os.path.join(OUT, f"ars_{tipo}.npz")
    rng = np.random.default_rng(args.seed)

    norm = Normalizzatore()
    storico = []
    it0 = 0
    migliore = -np.inf
    if args.resume:
        ck = np.load(args.resume)
        old_args = json.loads(str(ck['args']))
        if old_args.get('cmd', args.cmd) != args.cmd:
            ap.error('Il checkpoint usa un obiettivo diverso: avvia un nuovo job.')
        if 'rng_state' in ck and str(ck['rng_state']):
            rng.bit_generator.state = json.loads(str(ck['rng_state']))
        if 'migliore' in ck: migliore = float(ck['migliore'])
        assert str(ck["tipo"]) == tipo, "il checkpoint è di un altro tipo"
        theta = ck["theta"].copy()
        norm.n, norm.s1, norm.s2 = int(ck["norm_n"]), ck["norm_s1"].copy(), ck["norm_s2"].copy()
        it0 = int(ck["iterazione"])
        storico = [tuple(x) for x in ck["storico"]]
        print(f"Ripreso da {args.resume} (iterazione {it0})")
    elif tipo == "cpg":
        theta = cpgmod.gait_a_vettore(cpgmod.GAIT_PARTENZA)
    else:
        theta = theta_iniziale(rng)

    print(f"ARS V2-t | tipo={tipo} | parametri={theta.size} | N={args.dirs} b={args.top} "
          f"σ={sigma} α={lr} | episodio {args.episode}s | comandi '{args.cmd}' | "
          f"{args.workers} processi")
    print(f"{'iter':>5} {'tempo':>7} {'R medio(±)':>11} {'R valut.':>9} {'passi':>6} "
          f"{'cadute':>6} {'|θ|':>7}")
    eval_seeds = [10_000 + i for i in range(args.eval_eps)]
    t0 = time.time()
    ctx = mp.get_context("spawn")
    class Sequential:
        def map(self, fn, jobs): return list(map(fn, jobs))
    # Un worker: niente processo coordinatore o pool, un solo thread di calcolo.
    with (nullcontext(Sequential()) if args.workers == 1 else ctx.Pool(args.workers)) as pool:
        for it in range(it0 + 1, it0 + args.iters + 1):
            mean, std = norm.mean_std()
            dseeds = rng.integers(0, 2**31, size=args.dirs)
            eseeds = rng.integers(0, 2**31, size=args.dirs)
            lavori = [(tipo, theta, mean, std, int(ds), sigma, int(es), args.episode, args.cmd)
                      for ds, es in zip(dseeds, eseeds)]
            ris = pool.map(_lavoro, lavori)
            rp = np.array([a[0] for a, _ in ris]); rm = np.array([b[0] for _, b in ris])
            passi = sum(a[1] + b[1] for a, b in ris)
            if tipo == "mlp":
                for a, b in ris:
                    norm.aggiungi(*a[3]); norm.aggiungi(*b[3])
            # V2-t: tieni le b direzioni migliori
            ordine = np.argsort(-np.maximum(rp, rm))[:args.top]
            sr = np.concatenate([rp[ordine], rm[ordine]]).std() + 1e-8
            passo = np.zeros_like(theta)
            for k in ordine:
                delta = np.random.default_rng(int(dseeds[k])).standard_normal(theta.size)
                passo += (rp[k] - rm[k]) * delta
            theta = theta + lr / (args.top * sr) * passo
            if tipo == "cpg":
                theta = np.clip(theta, 0.0, 1.0)

            r_val = np.nan; cadute = ""
            if it % args.eval_every == 0 or it == it0 + args.iters or it == it0 + 1:
                m2, s2 = norm.mean_std()
                v = pool.map(_valuta, [(tipo, theta, m2, s2, s, args.episode, args.cmd)
                                       for s in eval_seeds])
                r_val = float(np.mean([x[0] for x in v]))
                cadute = f"{sum(x[2] for x in v)}/{len(v)}"
            r_med = float(np.mean(np.concatenate([rp, rm])))
            storico.append((it, time.time() - t0, r_med, r_val))
            if np.isfinite(r_val) and r_val > migliore:
                migliore = r_val
                salva(out_path.replace('.npz', '_migliore.npz'), tipo, theta, norm, it, args, storico, rng, migliore)
            if it % args.checkpoint_every == 0 or np.isfinite(r_val):
                salva(out_path, tipo, theta, norm, it, args, storico, rng, migliore)
            print(f"{it:5d} {time.time() - t0:6.0f}s {r_med:11.1f} "
                  f"{r_val:9.1f} {passi:6d} {cadute:>6} {np.linalg.norm(theta):7.2f}", flush=True)

    salva(out_path, tipo, theta, norm, it0 + args.iters, args, storico, rng, migliore)
    print(f"Checkpoint salvato: {out_path}")
    if tipo == "cpg":
        print(json.dumps(cpgmod.vettore_a_gait(theta), indent=2))


if __name__ == "__main__":
    main()
