| |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import sys |
| from collections import Counter, defaultdict |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| if str(PROJECT_ROOT) not in sys.path: |
| sys.path.insert(0, str(PROJECT_ROOT)) |
|
|
| from dovla_cil.data.datasets import CILDataset |
| from dovla_cil.eval.lattice_eval import _validation_group_ids |
| from dovla_cil.eval.maniskill_policy_rollout import ( |
| _nearest_retrieval_entries, |
| _numeric_action_values, |
| _select_action_chunk, |
| ) |
| from dovla_cil.models.dovla import ( |
| DoVLAConfig, |
| DoVLAModel, |
| load_model_state, |
| vectorize_toy_observation, |
| ) |
|
|
|
|
| def main(argv: list[str] | None = None) -> int: |
| parser = argparse.ArgumentParser( |
| description=( |
| "Export continuous BC targets from train-state counterfactual residual " |
| "retrieval scored by a trained DoVLA field." |
| ) |
| ) |
| parser.add_argument("--checkpoint", type=Path, required=True) |
| parser.add_argument("--dataset", type=Path, required=True) |
| parser.add_argument("--out", type=Path, required=True) |
| parser.add_argument("--device", default="auto") |
| parser.add_argument("--split", choices=("train", "val", "all"), default="all") |
| parser.add_argument("--retrieval-neighbors", type=int, default=1) |
| parser.add_argument( |
| "--retrieval-metric", |
| choices=("raw", "zscore", "task_relative"), |
| default="raw", |
| ) |
| parser.add_argument("--retrieval-residual-scale", type=float, default=0.35) |
| parser.add_argument( |
| "--exclude-types", |
| default="residual_random_negative,residual_wrong_direction,residual_near_miss", |
| help="Comma-separated residual candidate types to mask before field selection.", |
| ) |
| parser.add_argument( |
| "--no-leave-one-out", |
| action="store_true", |
| help="Allow a train target to retrieve residuals from its own source group.", |
| ) |
| parser.add_argument("--max-groups", type=int, default=None) |
| parser.add_argument("--clip-action-low", type=float, default=-1.0) |
| parser.add_argument("--clip-action-high", type=float, default=1.0) |
| parser.add_argument("--no-clip-actions", action="store_true") |
| args = parser.parse_args(argv) |
|
|
| if args.retrieval_neighbors <= 0: |
| parser.error("--retrieval-neighbors must be positive") |
| if args.retrieval_residual_scale < 0: |
| parser.error("--retrieval-residual-scale must be non-negative") |
| if args.max_groups is not None and args.max_groups <= 0: |
| parser.error("--max-groups must be positive when provided") |
| if args.clip_action_low >= args.clip_action_high: |
| parser.error("--clip-action-low must be smaller than --clip-action-high") |
|
|
| try: |
| import torch |
| except ImportError as exc: |
| raise ImportError("export_retrieval_residual_policy_targets.py requires torch") from exc |
|
|
| device = _resolve_device(args.device) |
| checkpoint = torch.load(args.checkpoint, map_location=device, weights_only=False) |
| model_config = DoVLAConfig(**checkpoint["model_config"]) |
| if model_config.observation_mode != "state": |
| raise ValueError("retrieval-residual target export currently supports state observations") |
| model = DoVLAModel(model_config).to(device) |
| load_model_state(model, checkpoint) |
| model.eval() |
|
|
| dataset = CILDataset(args.dataset) |
| trainer_config = checkpoint.get("trainer_config", {}) |
| val_ids = set( |
| _validation_group_ids( |
| dataset.group_ids, |
| val_fraction=float(trainer_config.get("val_fraction", 0.2)), |
| seed=int(trainer_config.get("seed", 0)), |
| ) |
| ) |
| train_ids = [group_id for group_id in dataset.group_ids if group_id not in val_ids] |
| if args.split == "train": |
| target_group_ids = list(train_ids) |
| elif args.split == "val": |
| target_group_ids = [group_id for group_id in dataset.group_ids if group_id in val_ids] |
| else: |
| target_group_ids = list(dataset.group_ids) |
| if args.max_groups is not None: |
| target_group_ids = target_group_ids[: args.max_groups] |
|
|
| bank = _build_residual_bank( |
| dataset, |
| train_ids, |
| obs_dim=model_config.obs_dim, |
| ) |
| excluded = {item.strip() for item in args.exclude_types.split(",") if item.strip()} |
| action_low = action_high = None |
| if not args.no_clip_actions: |
| action_low = torch.full( |
| (1, 1, model_config.action_dim), |
| float(args.clip_action_low), |
| dtype=torch.float32, |
| device=device, |
| ) |
| action_high = torch.full( |
| (1, 1, model_config.action_dim), |
| float(args.clip_action_high), |
| dtype=torch.float32, |
| device=device, |
| ) |
|
|
| targets: dict[str, dict[str, Any]] = {} |
| counts: Counter[str] = Counter() |
| source_counts: Counter[str] = Counter() |
| skipped: dict[str, str] = {} |
| with torch.no_grad(): |
| for group_id in target_group_ids: |
| records = dataset.get_group(group_id) |
| if not records: |
| skipped[group_id] = "empty_group" |
| continue |
| task_ids = {record.task_id for record in records} |
| if len(task_ids) != 1: |
| skipped[group_id] = "multi_task_group" |
| continue |
| task_id = next(iter(task_ids)) |
| entries = bank.get(task_id, []) |
| if args.no_leave_one_out: |
| candidates = entries |
| else: |
| candidates = [entry for entry in entries if entry[0] != group_id] |
| if not candidates: |
| candidates = entries |
| if not candidates: |
| skipped[group_id] = "no_train_bank_for_task" |
| continue |
|
|
| query = np.asarray( |
| vectorize_toy_observation( |
| records[0].observation_inline or {}, |
| obs_dim=model_config.obs_dim, |
| ), |
| dtype=np.float32, |
| ) |
| nearest = _nearest_retrieval_entries( |
| candidates, |
| query, |
| retrieval_neighbors=args.retrieval_neighbors, |
| retrieval_metric=args.retrieval_metric, |
| ) |
| source_group_ids: list[str] = [] |
| residuals: list[list[list[float]]] = [] |
| candidate_types: list[str] = [] |
| for source_group_id, _feature, source_residuals, source_types in nearest: |
| source_group_ids.append(source_group_id) |
| residuals.extend(source_residuals) |
| candidate_types.extend(source_types) |
|
|
| allowed = [candidate_type not in excluded for candidate_type in candidate_types] |
| if not any(allowed): |
| allowed = [True] * len(candidate_types) |
| obs = torch.tensor( |
| [ |
| vectorize_toy_observation( |
| records[0].observation_inline or {}, |
| obs_dim=model_config.obs_dim, |
| ) |
| ], |
| dtype=torch.float32, |
| device=device, |
| ) |
| action_residuals = torch.tensor( |
| [residuals], |
| dtype=torch.float32, |
| device=device, |
| ) |
| candidate_mask = torch.tensor([allowed], dtype=torch.bool, device=device) |
| selected, selected_index = _select_action_chunk( |
| model, |
| obs, |
| [records[0].instruction], |
| torch=torch, |
| selection_mode="retrieval_residual", |
| num_candidates=1, |
| candidate_sigma=0.0, |
| selection_seed=0, |
| retrieval_residual_scale=args.retrieval_residual_scale, |
| action_low=action_low, |
| action_high=action_high, |
| action_candidates=action_residuals, |
| candidate_mask=candidate_mask, |
| ) |
| index = int(selected_index[0]) |
| selected_type = candidate_types[index] if index < len(candidate_types) else "unknown" |
| action_values = selected[0].detach().cpu().tolist() |
| counts[selected_type] += 1 |
| for source_group_id in source_group_ids: |
| source_counts[source_group_id] += 1 |
| targets[group_id] = { |
| "action_values": action_values, |
| "selected_candidate_type": selected_type, |
| "candidate_source_group_id": ";".join(source_group_ids), |
| "task_id": task_id, |
| "retrieval_neighbors": args.retrieval_neighbors, |
| "retrieval_metric": args.retrieval_metric, |
| "retrieval_residual_scale": args.retrieval_residual_scale, |
| "excluded_candidate_types": sorted(excluded), |
| "leave_one_out": not args.no_leave_one_out, |
| } |
|
|
| payload = { |
| "target_type": "retrieval_residual_action_values", |
| "checkpoint": str(args.checkpoint), |
| "dataset": str(args.dataset), |
| "split": args.split, |
| "train_bank_groups": len(train_ids), |
| "num_groups": len(target_group_ids), |
| "num_targets": len(targets), |
| "num_skipped": len(skipped), |
| "retrieval_neighbors": args.retrieval_neighbors, |
| "retrieval_metric": args.retrieval_metric, |
| "retrieval_residual_scale": args.retrieval_residual_scale, |
| "excluded_candidate_types": sorted(excluded), |
| "leave_one_out": not args.no_leave_one_out, |
| "clip_actions": not args.no_clip_actions, |
| "clip_action_low": None if args.no_clip_actions else args.clip_action_low, |
| "clip_action_high": None if args.no_clip_actions else args.clip_action_high, |
| "selected_candidate_type_counts": dict(counts), |
| "num_source_groups_used": len(source_counts), |
| "skipped": skipped, |
| "targets": targets, |
| } |
| args.out.parent.mkdir(parents=True, exist_ok=True) |
| args.out.write_text(json.dumps(payload, indent=2) + "\n") |
| print(json.dumps({key: value for key, value in payload.items() if key != "targets"}, indent=2)) |
| print(f"Wrote {args.out}") |
| return 0 |
|
|
|
|
| def _build_residual_bank( |
| dataset: CILDataset, |
| group_ids: list[str], |
| *, |
| obs_dim: int, |
| ) -> dict[str, list[tuple[str, np.ndarray, list[list[list[float]]], list[str]]]]: |
| bank: dict[str, list[tuple[str, np.ndarray, list[list[list[float]]], list[str]]]] = ( |
| defaultdict(list) |
| ) |
| for group_id in group_ids: |
| records = dataset.get_group(group_id) |
| if not records: |
| continue |
| task_ids = {record.task_id for record in records} |
| if len(task_ids) != 1: |
| continue |
| anchor = next((record for record in records if record.candidate_type == "expert"), records[0]) |
| anchor_action = np.asarray(_numeric_action_values(anchor), dtype=np.float32) |
| residuals: list[list[list[float]]] = [np.zeros_like(anchor_action).tolist()] |
| candidate_types = ["policy_residual"] |
| for record in records: |
| if record.record_id == anchor.record_id: |
| continue |
| residual = np.asarray(_numeric_action_values(record), dtype=np.float32) - anchor_action |
| residuals.append(residual.tolist()) |
| candidate_types.append(f"residual_{record.candidate_type}") |
| feature = np.asarray( |
| vectorize_toy_observation(records[0].observation_inline or {}, obs_dim=obs_dim), |
| dtype=np.float32, |
| ) |
| bank[next(iter(task_ids))].append((group_id, feature, residuals, candidate_types)) |
| return bank |
|
|
|
|
| def _resolve_device(device: str) -> str: |
| if device != "auto": |
| return device |
| try: |
| import torch |
| except ImportError: |
| return "cpu" |
| return "cuda" if torch.cuda.is_available() else "cpu" |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|