#!/usr/bin/env python 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 # noqa: E402 from dovla_cil.eval.lattice_eval import _validation_group_ids # noqa: E402 from dovla_cil.eval.maniskill_policy_rollout import ( # noqa: E402 _nearest_retrieval_entries, _numeric_action_values, _select_action_chunk, ) from dovla_cil.models.dovla import ( # noqa: E402 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: # pragma: no cover 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: # pragma: no cover return "cpu" return "cuda" if torch.cuda.is_available() else "cpu" if __name__ == "__main__": raise SystemExit(main())