"""
Train the goal-conditioned action-chunking BC model for the demo2 arm.

Reuses the RX1 v2 architecture (models/act_model.py: ActionChunkPolicy +
Normalizer + save_checkpoint) — maps obs (init arm qpos + pick + place) → the
full pick-place action chunk (CHUNK, 6) = [5 arm ctrl, grasp]. The runtime plays
the predicted chunk open-loop (sidesteps BC covariate shift, like RX1).

CPU smoke-train here; GPU-scale on the Linux box for quality.

Run:  python demo2/train_arm.py --epochs 300  -> model/arm_act.pt
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path
import numpy as np
import torch

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE.parent))   # reach models/
from models.act_model import ActionChunkPolicy, Normalizer, save_checkpoint

DATA = HERE / "model" / "demos_arm.npz"
OUT = HERE / "model" / "arm_act.pt"


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--epochs", type=int, default=300)
    ap.add_argument("--lr", type=float, default=1e-3)
    ap.add_argument("--batch", type=int, default=64)
    ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
    args = ap.parse_args()

    z = np.load(DATA)
    obs = z["obs"].astype(np.float32)                 # (N, 11)
    chunk = z["chunk"].astype(np.float32)             # (N, CHUNK, 6)
    N, chunk_len, act_dim = chunk.shape
    obs_dim = obs.shape[1]
    flat = chunk.reshape(N, -1)                       # (N, CHUNK*6)
    print(f"[train] N={N} obs_dim={obs_dim} chunk={chunk_len}x{act_dim} device={args.device}")

    # normalise obs and (flattened) actions; reuse RX1 Normalizer
    norm = Normalizer.fit(obs, flat)
    obs_n = torch.tensor(norm.norm_obs(obs), dtype=torch.float32, device=args.device)
    act_n = torch.tensor(norm.norm_act(flat), dtype=torch.float32, device=args.device)

    model = ActionChunkPolicy(obs_dim, act_dim, chunk_len).to(args.device)
    opt = torch.optim.Adam(model.parameters(), lr=args.lr)
    lossf = torch.nn.SmoothL1Loss()

    for ep in range(args.epochs):
        perm = torch.randperm(N, device=args.device)
        tot = 0.0
        for i in range(0, N, args.batch):
            b = perm[i:i + args.batch]
            pred = model(obs_n[b]).reshape(len(b), -1)
            loss = lossf(pred, act_n[b])
            opt.zero_grad(); loss.backward(); opt.step()
            tot += loss.item() * len(b)
        if (ep + 1) % 50 == 0 or ep == 0:
            print(f"  epoch {ep+1:4d}  loss {tot/N:.5f}")

    save_checkpoint(str(OUT), model, norm)
    print(f"[train] saved -> {OUT.name}")


if __name__ == "__main__":
    main()
