vla / scripts /export_retrieval_residual_policy_targets.py
anhtld's picture
Auto-sync: 2026-06-28 22:04:20 (part 4)
ffada60 verified
Raw
History Blame Contribute Delete
12.1 kB
#!/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())