#!/usr/bin/env python3 """Train an A2C2 correction head on cached BEHAVIOR/OpenPI parquet data.""" from __future__ import annotations import argparse from dataclasses import asdict import json from pathlib import Path import random import sys import time import numpy as np import torch import torch.nn.functional as F from torch.utils.data import DataLoader SCRIPT_ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(SCRIPT_ROOT / "src")) from dataset import ( # noqa: E402 A2C2RandomSampleDataset, discover_episode_pairs, move_batch_to_device, pick_device, resolve_dataset_root, split_episode_pairs, ) from model import A2C2CorrectionHead, A2C2CorrectionHeadConfig # noqa: E402 def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--dataset-root", type=Path, required=True) parser.add_argument("--output-dir", type=Path, default=Path("a2c2/runs/task18")) parser.add_argument("--task-dir", default=None, help="Optional task directory filter, e.g. task-0018.") parser.add_argument("--steps", type=int, default=200_000) parser.add_argument("--batch-size", type=int, default=64) parser.add_argument("--num-workers", type=int, default=4) parser.add_argument("--samples-per-episode", type=int, default=512) parser.add_argument("--lr", type=float, default=1e-5) parser.add_argument("--weight-decay", type=float, default=1e-5) parser.add_argument("--grad-clip-norm", type=float, default=10.0) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--val-ratio", type=float, default=0.05) parser.add_argument("--max-episodes", type=int, default=None) parser.add_argument("--log-every", type=int, default=100) parser.add_argument("--save-every", type=int, default=10_000) parser.add_argument("--eval-every", type=int, default=0, help="Run validation every N steps. 0 disables validation.") parser.add_argument("--eval-samples", type=int, default=4096) parser.add_argument("--eval-batch-size", type=int, default=256) parser.add_argument("--device", default="auto") parser.add_argument("--dim-model", type=int, default=512) parser.add_argument("--n-heads", type=int, default=8) parser.add_argument("--n-encoder-layers", type=int, default=6) parser.add_argument("--dim-feedforward", type=int, default=2048) parser.add_argument("--dropout", type=float, default=0.1) parser.add_argument("--mlp-hidden-dim", type=int, default=1024) parser.add_argument( "--use-latent", dest="use_latent", action=argparse.BooleanOptionalAction, default=True, help="Use base-policy latent z during training. Pass --no-use-latent to train without latent.", ) parser.add_argument("--wandb", action="store_true", help="Enable Weights & Biases logging.") parser.add_argument("--wandb-project", default="a2c2") parser.add_argument("--wandb-entity", default=None) parser.add_argument("--wandb-run-name", default=None) parser.add_argument("--wandb-mode", default=None, choices=("online", "offline", "disabled")) return parser.parse_args() def save_checkpoint( output_dir: Path, model: A2C2CorrectionHead, optimizer: torch.optim.Optimizer, step: int, args: argparse.Namespace, ) -> Path: output_dir.mkdir(parents=True, exist_ok=True) path = output_dir / f"checkpoint_step_{step:06d}.pt" payload = { "step": step, "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "config": asdict(model.config), "args": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, } torch.save(payload, path) torch.save(payload, output_dir / "latest.pt") return path def init_wandb( args: argparse.Namespace, cfg: A2C2CorrectionHeadConfig, dataset_root: Path, train_episodes: int, val_episodes: int, num_parameters: int, ): if not args.wandb: return None try: import wandb except ImportError as exc: raise ImportError("wandb logging was requested. Install it with `pip install wandb`.") from exc run_config = { "dataset_root": str(dataset_root), "train_episodes": train_episodes, "val_episodes": val_episodes, "num_parameters": num_parameters, "model_config": asdict(cfg), "args": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, } return wandb.init( project=args.wandb_project, entity=args.wandb_entity, name=args.wandb_run_name, mode=args.wandb_mode, config=run_config, dir=str(args.output_dir), ) @torch.no_grad() def evaluate_model( model: A2C2CorrectionHead, val_pairs, cfg: A2C2CorrectionHeadConfig, device: torch.device, batch_size: int, num_samples: int, samples_per_episode: int, seed: int, ) -> dict[str, float]: if not val_pairs: return {} was_training = model.training model.eval() dataset = A2C2RandomSampleDataset( val_pairs, action_horizon=cfg.action_horizon, samples_per_episode=samples_per_episode, seed=seed, total_samples=num_samples, ) loader = DataLoader( dataset, batch_size=batch_size, num_workers=0, pin_memory=device.type == "cuda", ) total = 0 residual_mse_sum = 0.0 residual_mae_sum = 0.0 corrected_mse_sum = 0.0 base_mse_sum = 0.0 for batch in loader: batch = move_batch_to_device(batch, device) pred_delta = model( batch["observation_state"], batch["base_action"], batch["base_action_chunk"], batch["base_policy_z"], batch["time_feature"], batch["valid_action_mask"], ) target_delta = batch["target_delta"] base_action = batch["base_action"] expert_action = batch["expert_action"] corrected_action = base_action + pred_delta batch_size_actual = target_delta.shape[0] total += batch_size_actual residual_mse_sum += F.mse_loss(pred_delta, target_delta, reduction="sum").item() residual_mae_sum += F.l1_loss(pred_delta, target_delta, reduction="sum").item() corrected_mse_sum += F.mse_loss(corrected_action, expert_action, reduction="sum").item() base_mse_sum += F.mse_loss(base_action, expert_action, reduction="sum").item() if total >= num_samples: break if was_training: model.train() denom = max(total * cfg.action_dim, 1) return { "val/residual_mse": residual_mse_sum / denom, "val/residual_mae": residual_mae_sum / denom, "val/corrected_action_mse": corrected_mse_sum / denom, "val/base_action_mse": base_mse_sum / denom, "val/samples": float(total), } def main() -> None: args = parse_args() torch.manual_seed(args.seed) np.random.seed(args.seed) random.seed(args.seed) dataset_root = resolve_dataset_root(args.dataset_root) pairs = discover_episode_pairs(dataset_root, args.task_dir) train_pairs, val_pairs = split_episode_pairs(pairs, args.val_ratio, args.seed, args.max_episodes) print(f"Dataset root: {dataset_root}") print(f"Episodes: train={len(train_pairs)} val={len(val_pairs)}") cfg = A2C2CorrectionHeadConfig( use_base_policy_z=args.use_latent, dim_model=args.dim_model, n_heads=args.n_heads, n_encoder_layers=args.n_encoder_layers, dim_feedforward=args.dim_feedforward, dropout=args.dropout, mlp_hidden_dim=args.mlp_hidden_dim, ) device = pick_device(args.device) model = A2C2CorrectionHead(cfg).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay) num_parameters = sum(param.numel() for param in model.parameters()) train_dataset = A2C2RandomSampleDataset( train_pairs, action_horizon=cfg.action_horizon, samples_per_episode=args.samples_per_episode, seed=args.seed, ) train_loader = DataLoader( train_dataset, batch_size=args.batch_size, num_workers=args.num_workers, pin_memory=device.type == "cuda", ) train_iter = iter(train_loader) if args.eval_every > 0 and not val_pairs: print("WARNING: --eval-every was set, but validation split is empty. Validation will be skipped.", flush=True) args.output_dir.mkdir(parents=True, exist_ok=True) with (args.output_dir / "run_config.json").open("w", encoding="utf-8") as f: json.dump( { "dataset_root": str(dataset_root), "model_config": asdict(cfg), "args": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, }, f, indent=2, ) wandb_run = init_wandb( args=args, cfg=cfg, dataset_root=dataset_root, train_episodes=len(train_pairs), val_episodes=len(val_pairs), num_parameters=num_parameters, ) model.train() running_loss = 0.0 start = time.time() for step in range(1, args.steps + 1): batch = move_batch_to_device(next(train_iter), device) pred_delta = model( batch["observation_state"], batch["base_action"], batch["base_action_chunk"], batch["base_policy_z"], batch["time_feature"], batch["valid_action_mask"], ) loss = F.mse_loss(pred_delta, batch["target_delta"]) optimizer.zero_grad(set_to_none=True) loss.backward() grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip_norm) optimizer.step() loss_value = float(loss.detach().cpu()) running_loss += loss_value if wandb_run is not None: wandb_run.log( { "train/loss": loss_value, "train/grad_norm": float(grad_norm.detach().cpu()), "train/lr": args.lr, }, step=step, ) if step % args.log_every == 0: avg = running_loss / args.log_every elapsed = time.time() - start print(f"step={step} loss={avg:.6f} lr={args.lr:.2e} elapsed_s={elapsed:.1f}", flush=True) if wandb_run is not None: wandb_run.log( { "train/loss_avg": avg, "train/steps_per_second": args.log_every / max(elapsed, 1e-8), "train/elapsed_s_per_log_window": elapsed, }, step=step, ) running_loss = 0.0 start = time.time() if args.eval_every > 0 and val_pairs and step % args.eval_every == 0: metrics = evaluate_model( model=model, val_pairs=val_pairs, cfg=cfg, device=device, batch_size=args.eval_batch_size, num_samples=args.eval_samples, samples_per_episode=args.samples_per_episode, seed=args.seed + step, ) if metrics: print( "eval " f"step={step} " f"residual_mse={metrics['val/residual_mse']:.8f} " f"corrected_action_mse={metrics['val/corrected_action_mse']:.8f} " f"base_action_mse={metrics['val/base_action_mse']:.8f}", flush=True, ) if wandb_run is not None: wandb_run.log(metrics, step=step) if step % args.save_every == 0: path = save_checkpoint(args.output_dir, model, optimizer, step, args) print(f"saved {path}", flush=True) if wandb_run is not None: wandb_run.summary["latest_checkpoint"] = str(path) wandb_run.summary["latest_step"] = step path = save_checkpoint(args.output_dir, model, optimizer, args.steps, args) if wandb_run is not None: wandb_run.summary["final_checkpoint"] = str(path) wandb_run.summary["final_step"] = args.steps wandb_run.finish() print(f"training complete: {path}") if __name__ == "__main__": main()