worlddit / eval.py
Bidhan Roy
Release WorldDiT model and documentation
152b619
Raw
History Blame Contribute Delete
12.3 kB
#!/usr/bin/env python3
"""Evaluate a released WorldDiT checkpoint on LIBERO."""
from __future__ import annotations
import argparse
import copy
import contextlib
import json
import math
import os
import random
import subprocess
import sys
from collections import deque
from pathlib import Path
import clip
import numpy as np
import torch
from PIL import Image
from scipy.spatial.transform import Rotation
from torch.nn.parallel import DistributedDataParallel as DDP
from tqdm.auto import tqdm
from inference import ACTION_HORIZON, CONTEXT_STEPS, SUITES, load_model
def rank():
return int(os.environ.get("RANK", "0"))
def world_size():
return int(os.environ.get("WORLD_SIZE", "1"))
def launch(gpus: int):
if not 1 <= gpus <= torch.cuda.device_count():
raise ValueError(f"--gpus must be between 1 and {torch.cuda.device_count()}")
command = [
sys.executable,
"-m",
"torch.distributed.run",
"--standalone",
f"--nproc_per_node={gpus}",
str(Path(__file__).resolve()),
*sys.argv[1:],
]
environment = os.environ.copy()
environment.setdefault("OMP_NUM_THREADS", "1")
return subprocess.run(command, check=False, env=environment).returncode
def configure_libero(libero_root: Path, output: Path):
package = libero_root / "libero" / "libero"
required = (package / "bddl_files", package / "init_files", package / "assets")
if any(not path.is_dir() for path in required):
raise FileNotFoundError(f"invalid LIBERO checkout: {libero_root}")
config_dir = output / "libero_config"
if rank() == 0:
output.mkdir(parents=True, exist_ok=False)
config_dir.mkdir()
paths = {
"assets": str(package / "assets"),
"bddl_files": str(package / "bddl_files"),
"benchmark_root": str(package),
"datasets": str(package.parent / "datasets"),
"init_states": str(package / "init_files"),
}
(config_dir / "config.yaml").write_text(json.dumps(paths), encoding="utf-8")
torch.distributed.barrier()
os.environ["LIBERO_CONFIG_PATH"] = str(config_dir)
sys.path.insert(0, str(libero_root))
def orientation(quaternion):
quaternion = np.asarray(quaternion, dtype=np.float64).copy()
quaternion[3] = np.clip(quaternion[3], -1.0, 1.0)
denominator = np.sqrt(1.0 - quaternion[3] ** 2)
axisangle = (
np.zeros(3)
if math.isclose(denominator, 0.0)
else quaternion[:3] * 2.0 * math.acos(quaternion[3]) / denominator
)
return Rotation.from_euler("xyz", axisangle).as_euler("xyz")
def finish_action(action: torch.Tensor):
action = action.detach().cpu().numpy().copy()
action[-1] = 1.0 if action[-1] > 0.5 else -1.0
return action
class PolicyRunner:
def __init__(
self, model, temperature: float, execution_horizon: int, max_steps: int
):
self.model = model
self.processor = getattr(model, "module", model).image_processor
self.temperature = temperature
self.execution_horizon = execution_horizon
self.max_steps = max_steps
def reset(self):
self.primary = deque(maxlen=CONTEXT_STEPS)
self.wrist = deque(maxlen=CONTEXT_STEPS)
self.state = deque(maxlen=CONTEXT_STEPS)
self.pending = deque()
self.predictions = torch.zeros(
self.max_steps, self.max_steps + ACTION_HORIZON, 7, device="cuda"
)
self.valid = torch.zeros(
self.max_steps,
self.max_steps + ACTION_HORIZON,
dtype=torch.bool,
device="cuda",
)
def observe(self, observation):
primary = Image.fromarray(observation["agentview_image"][::-1])
wrist = Image.fromarray(observation["robot0_eye_in_hand_image"])
self.primary.append(self.processor(primary).unsqueeze(0).unsqueeze(0))
self.wrist.append(self.processor(wrist).unsqueeze(0).unsqueeze(0))
state = np.concatenate(
(
observation["robot0_eef_pos"],
orientation(observation["robot0_eef_quat"]),
observation["robot0_gripper_qpos"],
)
)
self.state.append(torch.from_numpy(state).float().view(1, 1, -1))
def prefill(self, observations):
for observation in observations[-CONTEXT_STEPS:-1]:
self.observe(observation)
def ensemble(self, chunk: torch.Tensor, timestep: int):
self.predictions[timestep, timestep : timestep + ACTION_HORIZON] = chunk
self.valid[timestep, timestep : timestep + ACTION_HORIZON] = True
selected = []
for target in range(timestep, timestep + self.execution_horizon):
mask = self.valid[: timestep + 1, target]
actions = self.predictions[: timestep + 1, target][mask]
weights = np.exp(-self.temperature * np.arange(len(actions)))
weights = torch.as_tensor(weights / weights.sum(), device="cuda").unsqueeze(
1
)
selected.append(finish_action((actions * weights).sum(0)))
self.pending.extend(selected[1:])
return selected[0]
@torch.inference_mode()
def act(self, observation, instruction: str, timestep: int):
self.observe(observation)
if self.pending:
return self.pending.popleft()
if len(self.primary) != CONTEXT_STEPS:
raise RuntimeError("evaluation requires three real context observations")
primary = torch.cat(tuple(self.primary), dim=1).cuda()
wrist = torch.cat(tuple(self.wrist), dim=1).cuda()
state = torch.cat(tuple(self.state), dim=1).cuda()
text = (
clip.tokenize([instruction] * CONTEXT_STEPS, truncate=True)
.view(1, CONTEXT_STEPS, -1)
.cuda()
)
chunk = self.model(primary, wrist, state, text)[0, -1]
return self.ensemble(chunk, timestep)
def evaluate(args, model):
with (
open(os.devnull, "w") as quiet,
contextlib.redirect_stdout(quiet),
contextlib.redirect_stderr(quiet),
):
from libero.libero import benchmark
from libero.libero.envs import OffScreenRenderEnv
suite = benchmark.get_benchmark_dict()[args.suite]()
runner = PolicyRunner(
model, args.temperature, args.execution_horizon, args.max_steps
)
total = args.tasks * args.episodes
assigned = list(range(total))[rank() :: world_size()]
local_results = []
progress = tqdm(
assigned,
desc=f"GPU {rank()}",
position=rank(),
dynamic_ncols=True,
leave=True,
)
for evaluation_id in progress:
task_id, episode_index = divmod(evaluation_id, args.episodes)
episode_id = args.episode_offset + episode_index
task = suite.get_task(task_id)
bddl = (
args.libero_path
/ "libero"
/ "libero"
/ "bddl_files"
/ task.problem_folder
/ task.bddl_file
)
environment = OffScreenRenderEnv(
bddl_file_name=str(bddl),
camera_heights=128,
camera_widths=128,
render_gpu_device_id=int(os.environ["LOCAL_RANK"]),
)
try:
environment.reset()
environment.seed(66)
initial_states = torch.load(
args.libero_path
/ "libero"
/ "libero"
/ "init_files"
/ task.problem_folder
/ task.init_states_file,
weights_only=False,
)
if episode_id >= len(initial_states):
raise IndexError(
f"episode {episode_id} is unavailable for task {task_id}"
)
observation = environment.set_init_state(initial_states[episode_id])
warmup = []
for _ in range(5):
observation, _, _, _ = environment.step(np.zeros(7))
warmup.append(copy.deepcopy(observation))
runner.reset()
runner.prefill(warmup)
observation = warmup[-1]
success = 0
for steps in range(1, args.max_steps + 1):
action = runner.act(observation, task.language, steps - 1)
observation, _, done, _ = environment.step(action)
if done:
success = 1
break
local_results.append((evaluation_id, task_id, episode_id, success, steps))
progress.set_postfix(successes=sum(item[3] for item in local_results))
finally:
environment.close()
gathered = [None] * world_size() if rank() == 0 else None
torch.distributed.gather_object(local_results, gathered, dst=0)
if rank() != 0:
return
results = sorted((item for group in gathered for item in group), key=lambda x: x[0])
per_task = []
print()
for task_id in range(args.tasks):
values = [item[3] for item in results if item[1] == task_id]
rate = float(np.mean(values))
per_task.append(rate)
print(f"Task {task_id}: {sum(values)}/{len(values)} ({rate:.1%})")
successes = sum(item[3] for item in results)
print(f"Overall: {successes}/{len(results)} ({successes / len(results):.1%})")
report = {
"suite": args.suite,
"gpus": world_size(),
"episodes": len(results),
"successes": successes,
"success_rate": successes / len(results),
"per_task_success_rate": per_task,
"results": [
{"task": item[1], "episode": item[2], "success": item[3], "steps": item[4]}
for item in results
],
}
(args.output_dir / "results.json").write_text(
json.dumps(report, indent=2) + "\n", encoding="utf-8"
)
def parser():
value = argparse.ArgumentParser(description=__doc__)
value.add_argument("--suite", required=True, choices=SUITES)
value.add_argument("--gpus", type=int, default=1)
value.add_argument(
"--model-root", type=Path, default=Path(__file__).resolve().parent
)
value.add_argument(
"--libero-path", type=Path, default=Path("~/LIBERO").expanduser()
)
value.add_argument("--output-dir", type=Path, required=True)
value.add_argument("--tasks", type=int, default=10)
value.add_argument("--episodes", type=int, default=50)
value.add_argument("--episode-offset", type=int, default=0)
value.add_argument("--max-steps", type=int, default=600)
value.add_argument("--execution-horizon", type=int, choices=(1, 3), default=3)
value.add_argument("--temperature", type=float, default=0.01)
return value
def main():
args = parser().parse_args()
if "RANK" not in os.environ:
return launch(args.gpus)
if args.gpus != world_size():
raise ValueError(
f"--gpus={args.gpus} but torchrun started {world_size()} workers"
)
args.model_root = args.model_root.expanduser().resolve()
args.libero_path = args.libero_path.expanduser().resolve()
args.output_dir = args.output_dir.expanduser().resolve()
os.environ.update(MUJOCO_GL="egl", PYOPENGL_PLATFORM="egl")
os.environ.setdefault("NCCL_IB_DISABLE", "1")
os.environ.setdefault("NCCL_P2P_DISABLE", "1")
os.environ.setdefault("NCCL_CUMEM_ENABLE", "0")
os.environ.setdefault("TORCH_NCCL_BLOCKING_WAIT", "1")
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
torch.distributed.init_process_group(
backend="nccl", device_id=torch.device("cuda", int(os.environ["LOCAL_RANK"]))
)
try:
configure_libero(args.libero_path, args.output_dir)
model = load_model(args.model_root, args.suite, torch.device("cuda"))
seed = 66 + rank()
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
model = DDP(
model,
device_ids=[int(os.environ["LOCAL_RANK"])],
find_unused_parameters=True,
)
model.eval()
evaluate(args, model)
torch.distributed.barrier()
finally:
torch.distributed.destroy_process_group()
return 0
if __name__ == "__main__":
raise SystemExit(main())