#!/usr/bin/env python3
"""
train_act.py  (v2) — behaviour-clone the learned arms+grasp action model.

Builds action-CHUNK targets from the demonstrations (each obs_t -> the next
`chunk_len` expert actions, respecting episode boundaries), normalises, and
fits models/act_model.py with MSE. Runs on GPU when available, CPU otherwise.

    python scripts/train_act.py --demos data/demos.npz --out models/ckpt/act.pt \
        --epochs 200 --chunk 8

On the RTX 3060 a few hundred epochs over a few thousand demos takes minutes.
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path

import numpy as np
import torch
from torch.utils.data import DataLoader, TensorDataset

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from models.act_model import ActionChunkPolicy, Normalizer, save_checkpoint  # noqa: E402


def build_chunks(obs, act, ep_len, chunk):
    """obs_t -> next `chunk` actions, padding the tail of each episode by repeat."""
    O, A = [], []
    i = 0
    for n in ep_len:
        eo, ea = obs[i:i + n], act[i:i + n]
        for t in range(n):
            ch = ea[t:t + chunk]
            if len(ch) < chunk:
                ch = np.concatenate([ch, np.repeat(ch[-1:], chunk - len(ch), 0)])
            O.append(eo[t]); A.append(ch)
        i += n
    return np.asarray(O, np.float32), np.asarray(A, np.float32)


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--demos", default="data/demos.npz")
    p.add_argument("--out", default="models/ckpt/act.pt")
    p.add_argument("--epochs", type=int, default=200)
    p.add_argument("--chunk", type=int, default=8)
    p.add_argument("--batch", type=int, default=256)
    p.add_argument("--lr", type=float, default=3e-4)
    p.add_argument("--hidden", type=int, default=512)
    p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
    a = p.parse_args()

    d = np.load(a.demos)
    obs, act, ep_len = d["obs"], d["act"], d["ep_len"]
    print(f"loaded {obs.shape[0]} transitions, {len(ep_len)} episodes, "
          f"obs_dim={obs.shape[1]} act_dim={act.shape[1]}  device={a.device}")

    O, A = build_chunks(obs, act, ep_len, a.chunk)             # (N,86) (N,chunk,15)
    norm = Normalizer.fit(O, A.reshape(-1, A.shape[-1]))
    On = norm.norm_obs(O)
    An = norm.norm_act(A.reshape(-1, A.shape[-1])).reshape(A.shape)

    ds = TensorDataset(torch.from_numpy(On).float(), torch.from_numpy(An).float())
    dl = DataLoader(ds, batch_size=a.batch, shuffle=True, drop_last=False)

    model = ActionChunkPolicy(obs.shape[1], act.shape[1], a.chunk, a.hidden).to(a.device)
    opt = torch.optim.AdamW(model.parameters(), lr=a.lr)
    lossf = torch.nn.MSELoss()

    model.train()
    for ep in range(a.epochs):
        tot = 0.0
        for ob, ac in dl:
            ob, ac = ob.to(a.device), ac.to(a.device)
            opt.zero_grad()
            loss = lossf(model(ob), ac)
            loss.backward(); opt.step()
            tot += loss.item() * ob.shape[0]
        if (ep + 1) % max(1, a.epochs // 10) == 0 or ep == 0:
            print(f"  epoch {ep+1:4d}/{a.epochs}   mse={tot/len(ds):.5f}")

    Path(a.out).parent.mkdir(parents=True, exist_ok=True)
    save_checkpoint(a.out, model, norm)
    print(f"saved checkpoint -> {a.out}")


if __name__ == "__main__":
    main()
