| from __future__ import annotations |
|
|
| import argparse |
| import json |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
| import torch |
| from omegaconf import OmegaConf |
|
|
| from eval.run_reveal_benchmark import load_model, _resolve_checkpoint_from_config |
| from sim_reveal.dataset import collect_teacher_dataset, save_teacher_dataset |
| from sim_reveal.procedural_envs import render_views_from_state |
| from train.dataset_build_utils import dataset_version_with_suffix, output_dataset_path |
|
|
|
|
| def _render_history( |
| proxy_name: str, |
| history_render_states: list[dict[str, Any]], |
| resolution: int, |
| ) -> tuple[list[np.ndarray], list[np.ndarray], list[np.ndarray]]: |
| history_images: list[np.ndarray] = [] |
| history_depths: list[np.ndarray] = [] |
| history_depth_valid: list[np.ndarray] = [] |
| for render_state in history_render_states: |
| rendered = render_views_from_state( |
| proxy_name=proxy_name, |
| render_state=render_state, |
| resolution=resolution, |
| include_depth=True, |
| ) |
| history_images.append( |
| np.stack([rendered["front"], rendered["wrist_left"], rendered["wrist_right"]], axis=0).astype(np.uint8) |
| ) |
| history_depths.append( |
| np.stack([rendered["front_depth"], rendered["wrist_left_depth"], rendered["wrist_right_depth"]], axis=0)[:, None, :, :].astype(np.float32) |
| ) |
| history_depth_valid.append( |
| np.stack( |
| [rendered["front_depth_valid"], rendered["wrist_left_depth_valid"], rendered["wrist_right_depth_valid"]], |
| axis=0, |
| )[:, None, :, :].astype(np.float32) |
| ) |
| return history_images, history_depths, history_depth_valid |
|
|
|
|
| def _prepare_model_inputs( |
| observation: dict[str, Any], |
| sample: dict[str, Any], |
| device: torch.device, |
| resolution: int, |
| ) -> dict[str, Any]: |
| history_render_states = list(sample.get("history_render_states", [])) |
| history_images, history_depths, history_depth_valid = _render_history( |
| proxy_name=str(sample["proxy_name"]), |
| history_render_states=history_render_states, |
| resolution=resolution, |
| ) |
| if history_images: |
| history_images_tensor = torch.from_numpy(np.stack(history_images, axis=0)).permute(0, 1, 4, 2, 3).unsqueeze(0).float() / 255.0 |
| history_depths_tensor = torch.from_numpy(np.stack(history_depths, axis=0)).unsqueeze(0).float() |
| history_depth_valid_tensor = torch.from_numpy(np.stack(history_depth_valid, axis=0)).unsqueeze(0).float() |
| else: |
| history_images_tensor = torch.zeros((1, 0, 3, 3, resolution, resolution), dtype=torch.float32) |
| history_depths_tensor = torch.zeros((1, 0, 3, 1, resolution, resolution), dtype=torch.float32) |
| history_depth_valid_tensor = torch.zeros_like(history_depths_tensor) |
| proprio_dim = observation["proprio"].shape[0] |
| return { |
| "images": torch.from_numpy(observation["images"]).permute(0, 3, 1, 2).unsqueeze(0).float().to(device) / 255.0, |
| "depths": torch.from_numpy(observation["depths"]).unsqueeze(0).float().to(device), |
| "depth_valid": torch.from_numpy(observation["depth_valid"]).unsqueeze(0).float().to(device), |
| "camera_intrinsics": torch.from_numpy(observation["camera_intrinsics"]).unsqueeze(0).float().to(device), |
| "camera_extrinsics": torch.from_numpy(observation["camera_extrinsics"]).unsqueeze(0).float().to(device), |
| "proprio": torch.from_numpy(observation["proprio"]).unsqueeze(0).float().to(device), |
| "texts": [str(observation["text"])], |
| "task_names": [str(sample["task_name"])], |
| "task_ids": torch.as_tensor([int(sample["task_id"])], dtype=torch.long, device=device), |
| "history_images": history_images_tensor.to(device), |
| "history_depths": history_depths_tensor.to(device), |
| "history_depth_valid": history_depth_valid_tensor.to(device), |
| "history_camera_intrinsics": torch.from_numpy( |
| sample.get("history_camera_intrinsics", np.zeros((0, 3, 3, 3), dtype=np.float32)) |
| ).unsqueeze(0).float().to(device), |
| "history_camera_extrinsics": torch.from_numpy( |
| sample.get("history_camera_extrinsics", np.zeros((0, 3, 4, 4), dtype=np.float32)) |
| ).unsqueeze(0).float().to(device), |
| "history_camera_valid_mask": torch.from_numpy( |
| sample.get("history_camera_valid_mask", np.zeros((0, 3), dtype=np.float32)) |
| ).unsqueeze(0).float().to(device), |
| "history_proprio": torch.from_numpy( |
| sample.get("history_proprio", np.zeros((0, proprio_dim), dtype=np.float32)) |
| ).unsqueeze(0).float().to(device), |
| "history_actions": torch.from_numpy( |
| sample.get("history_actions", np.zeros((0, sample["action_chunk"].shape[-1]), dtype=np.float32)) |
| ).unsqueeze(0).float().to(device), |
| } |
|
|
|
|
| def _proposal_target_builder(model: torch.nn.Module, device: torch.device, resolution: int): |
| def _build(env: Any, observation: dict[str, Any], sample: dict[str, Any]) -> dict[str, Any]: |
| with torch.inference_mode(): |
| outputs = model( |
| **_prepare_model_inputs(observation, sample, device, resolution), |
| plan=False, |
| use_planner=False, |
| use_world_model=False, |
| use_proposal_candidates=True, |
| ) |
| proposal_candidates = outputs["proposal_candidates"][0].detach().float().cpu().numpy().astype(np.float32) |
| outcomes = [env.evaluate_action_chunk(candidate, rollout_horizon=env.rollout_horizon) for candidate in proposal_candidates] |
| proposal_target_retrieval_success = np.asarray([item["retrieval_success"] for item in outcomes], dtype=np.float32) |
| proposal_target_risk = np.clip( |
| np.asarray([item["final_disturbance_cost"] + item["reocclusion_rate"] for item in outcomes], dtype=np.float32), |
| 0.0, |
| 1.0, |
| ).astype(np.float32) |
| proposal_target_utility = np.asarray( |
| [float(env.candidate_outcome_utility(item)) for item in outcomes], |
| dtype=np.float32, |
| ) |
| return { |
| "proposal_target_action_chunks": proposal_candidates, |
| "proposal_target_retrieval_success": proposal_target_retrieval_success, |
| "proposal_target_risk": proposal_target_risk, |
| "proposal_target_utility": proposal_target_utility, |
| "proposal_target_mode_names": list(outputs.get("proposal_mode_names", [["unknown"]])[0]), |
| } |
|
|
| return _build |
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", required=True) |
| parser.add_argument("--checkpoint", default=None) |
| parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") |
| parser.add_argument("--train-output", default=None) |
| parser.add_argument("--val-output", default=None) |
| parser.add_argument("--dataset-suffix", default="selector_align") |
| args = parser.parse_args() |
|
|
| cfg = OmegaConf.load(args.config) |
| checkpoint_path = Path(args.checkpoint) if args.checkpoint else _resolve_checkpoint_from_config(args.config) |
| device = torch.device(args.device) |
| model, _ = load_model(checkpoint_path, device) |
| model.eval() |
|
|
| resolution = int(cfg.data.resolution) |
| builder = _proposal_target_builder(model, device, resolution) |
| dataset_version = dataset_version_with_suffix( |
| str(cfg.data.get("dataset_version", "reveal_proxy_v6")), |
| args.dataset_suffix, |
| ) |
| train_output = Path(args.train_output) if args.train_output else output_dataset_path(cfg.data.train_dataset_path, args.dataset_suffix) |
| val_output = Path(args.val_output) if args.val_output else output_dataset_path(cfg.data.val_dataset_path, args.dataset_suffix) |
|
|
| bundles: dict[str, dict[str, Any]] = {} |
| for split, episodes_per_proxy, seed_offset, output_path in ( |
| ("train", int(cfg.data.train_episodes_per_proxy), 0, train_output), |
| ("val", int(cfg.data.val_episodes_per_proxy), 10_000, val_output), |
| ): |
| bundle = collect_teacher_dataset( |
| proxy_names=OmegaConf.to_container(cfg.data.proxies, resolve=True), |
| episodes_per_proxy=episodes_per_proxy, |
| resolution=resolution, |
| seed=int(cfg.data.seed) + seed_offset, |
| chunk_horizon=int(cfg.data.chunk_horizon), |
| rollout_horizon=int(cfg.data.rollout_horizon), |
| history_steps=int(cfg.data.get("history_steps", 2)), |
| planner_candidates=int(cfg.data.get("planner_candidates", 4)), |
| dataset_version=dataset_version, |
| proposal_target_builder=builder, |
| ) |
| save_teacher_dataset(output_path, bundle) |
| bundles[split] = { |
| "output_path": str(output_path), |
| "samples": len(bundle["samples"]), |
| "dataset_version": dataset_version, |
| } |
| print(json.dumps({"phase": "dataset_saved", "split": split, **bundles[split]}), flush=True) |
|
|
| summary = { |
| "checkpoint": str(checkpoint_path), |
| "device": str(device), |
| "dataset_suffix": args.dataset_suffix, |
| "train": bundles["train"], |
| "val": bundles["val"], |
| } |
| summary_path = train_output.parent / f"proposal_dataset_build_{args.dataset_suffix}.json" |
| summary_path.write_text(json.dumps(summary, indent=2), encoding="utf-8") |
| print(json.dumps(summary, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|