| |
| """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 ( |
| A2C2RandomSampleDataset, |
| discover_episode_pairs, |
| move_batch_to_device, |
| pick_device, |
| resolve_dataset_root, |
| split_episode_pairs, |
| ) |
| from model import A2C2CorrectionHead, A2C2CorrectionHeadConfig |
|
|
|
|
| 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() |
|
|