a2c2 / scripts /train.py
dennis96's picture
Upload folder using huggingface_hub
07f85c4 verified
Raw
History Blame Contribute Delete
12.7 kB
#!/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()