Files
roboimi/roboimi/demos/vla_scripts/eval_vla_parallel_tmp.py
T

1756 lines
65 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
VLA 策略评估脚本(简化版)
该脚本使用 agent 内置的队列管理来评估训练好的 VLA 策略。
无需单独的评估器类 - agent 处理一切!
使用方法:
python roboimi/demos/eval_vla_simple.py
python roboimi/demos/eval_vla_simple.py eval.ckpt_path=checkpoints/vla_model_final.pt
python roboimi/demos/eval_vla_simple.py eval.ckpt_path=checkpoints/vla_model_best.pt
"""
import sys
import os
import json
import logging
import time
import queue
import concurrent.futures
import multiprocessing
import torch
import numpy as np
import hydra
from pathlib import Path
from collections import deque
from typing import Any, Dict, Optional
from tqdm import tqdm
from omegaconf import DictConfig, OmegaConf
from hydra.utils import instantiate
from einops import rearrange
from roboimi.envs.double_pos_ctrl_env import make_sim_env
from roboimi.utils.act_ex_utils import sample_transfer_pose
from roboimi.vla.eval_utils import execute_policy_action
sys.path.append(os.getcwd())
log = logging.getLogger(__name__)
if not OmegaConf.has_resolver("len"):
OmegaConf.register_new_resolver("len", lambda x: len(x))
def load_checkpoint(
ckpt_path: str,
agent_cfg: DictConfig,
device: str = 'cuda'
) -> torch.nn.Module:
"""
从检查点加载训练好的 VLA 模型,使用 Hydra agent 配置。
Args:
ckpt_path: 检查点文件路径 (.pt)
agent_cfg: Hydra agent 配置,用于实例化
device: 加载模型的设备
Returns:
加载后的 VLAAgent 模型
"""
from pathlib import Path as PathLib
ckpt_path = PathLib(ckpt_path).absolute()
if not ckpt_path.exists():
raise FileNotFoundError(f"检查点未找到: {ckpt_path}")
log.info(f"{ckpt_path} 加载检查点")
checkpoint = torch.load(ckpt_path, map_location=device, weights_only=False)
log.info(f"检查点键值: {checkpoint.keys()}")
# 加载数据集统计信息用于归一化
stats = checkpoint.get('dataset_stats', None)
# 使用数据集统计信息从 Hydra 配置实例化 agent
log.info("从配置实例化 agent...")
agent = instantiate(agent_cfg, dataset_stats=stats)
# 加载模型状态
agent.load_state_dict(checkpoint['model_state_dict'])
log.info(f"✅ 模型状态已加载 (步数: {checkpoint.get('step', 'unknown')})")
if stats is not None:
log.info(f"✅ 数据集统计信息已加载 (归一化: {stats.get('normalization_type', 'gaussian')})")
else:
# 后备方案:尝试从外部 JSON 文件加载(兼容旧检查点)
stats_path = ckpt_path.parent / 'dataset_stats.json'
if stats_path.exists():
with open(stats_path, 'r') as f:
stats = json.load(f)
log.info("✅ 数据集统计信息已从外部 JSON 加载(旧版本兼容)")
else:
log.warning("⚠️ 未找到数据集统计信息。动作将无法反归一化!")
agent.eval()
agent.to(device)
log.info(f"✅ 模型已成功加载到 {device}")
return agent, stats
def prepare_observation(obs: Dict, camera_names: list) -> Dict:
"""
将环境观测转换为 agent 格式。
Args:
obs: 环境观测字典,包含图像和 qpos
camera_names: 摄像头名称列表
Returns:
agent 格式的观测字典
"""
import cv2
# 转换图像: numpy -> tensor, HWC -> CHW
images = {}
for cam_name in camera_names:
img = obs['images'][cam_name]
# Resize 到 224x224(与训练时一致)
img = cv2.resize(img, (224, 224), interpolation=cv2.INTER_LINEAR)
img = rearrange(img, 'h w c -> c h w')
img = torch.from_numpy(img / 255.0).float()
images[cam_name] = img
# 转换 qpos: numpy -> tensor
qpos = torch.from_numpy(obs['qpos']).float()
return {'qpos': qpos, 'images': images}
def _resolve_policy_camera_names(cfg: DictConfig) -> list[str]:
agent_cfg = cfg.agent
eval_cfg = cfg.eval
camera_names = agent_cfg.get('camera_names', None)
if camera_names is not None:
return list(camera_names)
agent_target = str(agent_cfg.get('_target_', ''))
if agent_target.endswith('VLAAgentGr00tDiT'):
return list(eval_cfg.camera_names)
return sorted(eval_cfg.camera_names)
def _new_local_policy_queues(obs_horizon: int) -> dict[str, deque]:
return {
'qpos': deque(maxlen=int(obs_horizon)),
'images': deque(maxlen=int(obs_horizon)),
'action': deque(),
}
def _populate_local_policy_queues(
queues: dict[str, deque],
observation: Dict[str, torch.Tensor],
) -> None:
if 'qpos' in observation:
queues['qpos'].append(observation['qpos'].detach().clone())
if 'images' in observation:
queues['images'].append({
camera_name: image.detach().clone()
for camera_name, image in observation['images'].items()
})
def _prepare_local_policy_batch(
queues: dict[str, deque],
obs_horizon: int,
camera_names: list[str],
) -> Dict[str, torch.Tensor]:
qpos_list = list(queues['qpos'])
if not qpos_list:
raise ValueError('observation queue is empty.')
while len(qpos_list) < int(obs_horizon):
qpos_list.append(qpos_list[-1])
batch_qpos = torch.stack(qpos_list, dim=0).unsqueeze(0)
images_list = list(queues['images'])
if not images_list:
raise ValueError('image queue is empty.')
while len(images_list) < int(obs_horizon):
images_list.append(images_list[-1])
ordered_camera_names = list(camera_names) if camera_names else sorted(images_list[0].keys())
batch_images = {
camera_name: torch.stack(
[image_history[camera_name] for image_history in images_list],
dim=0,
).unsqueeze(0)
for camera_name in ordered_camera_names
}
return {'qpos': batch_qpos, 'images': batch_images}
def _enqueue_predicted_actions(
queues: dict[str, deque],
predicted_actions: Any,
obs_horizon: int,
num_action_steps: int,
) -> None:
if isinstance(predicted_actions, np.ndarray):
predicted_actions = torch.from_numpy(predicted_actions)
if predicted_actions.ndim == 2:
predicted_actions = predicted_actions.unsqueeze(0)
start = int(obs_horizon) - 1
end = start + int(num_action_steps)
executable_actions = predicted_actions[:, start:end]
for action_index in range(executable_actions.shape[1]):
queues['action'].append(
executable_actions[:, action_index].squeeze(0).detach().cpu().clone()
)
def _serialize_policy_batch(batch: Dict[str, torch.Tensor]) -> dict[str, Any]:
return {
'qpos': batch['qpos'].detach().cpu().numpy().astype(np.float32, copy=True),
'images': {
camera_name: image.detach().cpu().numpy().astype(np.float32, copy=True)
for camera_name, image in batch['images'].items()
},
}
def _deserialize_policy_batch(batch: dict[str, Any], device: str) -> Dict[str, torch.Tensor]:
return {
'qpos': torch.as_tensor(batch['qpos'], dtype=torch.float32, device=device),
'images': {
camera_name: torch.as_tensor(image, dtype=torch.float32, device=device)
for camera_name, image in batch['images'].items()
},
}
class _LocalPolicyRunner:
def __init__(self, agent: torch.nn.Module):
self.agent = agent
self.uses_local_model = True
def reset(self):
self.agent.reset()
def select_action(
self,
observation: Dict[str, torch.Tensor],
*,
episode_index: int,
timestep: int,
) -> tuple[Any, bool]:
del episode_index, timestep
action_queue = getattr(self.agent, '_queues', {}).get('action', None)
model_inference_triggered = len(action_queue) == 0 if action_queue is not None else True
action = self.agent.select_action(observation)
return action, bool(model_inference_triggered)
class _RemotePolicyRunner:
def __init__(
self,
*,
worker_index: int,
server_index: int,
request_queue,
response_queue,
camera_names: list[str],
obs_horizon: int,
num_action_steps: int,
response_timeout_s: float = 30.0,
):
self.worker_index = int(worker_index)
self.server_index = int(server_index)
self.request_queue = request_queue
self.response_queue = response_queue
self.camera_names = list(camera_names)
self.obs_horizon = int(obs_horizon)
self.num_action_steps = int(num_action_steps)
self.response_timeout_s = float(response_timeout_s)
self.local_queues = _new_local_policy_queues(self.obs_horizon)
self.uses_local_model = False
def reset(self):
self.local_queues = _new_local_policy_queues(self.obs_horizon)
def select_action(
self,
observation: Dict[str, torch.Tensor],
*,
episode_index: int,
timestep: int,
) -> tuple[torch.Tensor, bool]:
_populate_local_policy_queues(self.local_queues, observation)
model_inference_triggered = len(self.local_queues['action']) == 0
if model_inference_triggered:
batch = _prepare_local_policy_batch(
self.local_queues,
obs_horizon=self.obs_horizon,
camera_names=self.camera_names,
)
self.request_queue.put({
'type': 'predict_chunk',
'worker_index': self.worker_index,
'server_index': self.server_index,
'episode_index': int(episode_index),
'timestep': int(timestep),
'batch': _serialize_policy_batch(batch),
})
try:
response = self.response_queue.get(timeout=self.response_timeout_s)
except queue.Empty as exc:
raise RuntimeError(
f'worker {self.worker_index} timed out waiting for inference server {self.server_index}'
) from exc
if response.get('type') != 'predict_chunk_result':
raise RuntimeError(
f'worker {self.worker_index} received unexpected inference response: '
f'{response.get("type")}'
)
_enqueue_predicted_actions(
self.local_queues,
predicted_actions=response['actions'],
obs_horizon=self.obs_horizon,
num_action_steps=self.num_action_steps,
)
if not self.local_queues['action']:
raise RuntimeError(f'worker {self.worker_index} received no executable action from server')
return self.local_queues['action'].popleft(), bool(model_inference_triggered)
def _to_numpy_action(action: Any) -> np.ndarray:
if isinstance(action, torch.Tensor):
return action.detach().cpu().numpy().astype(np.float32, copy=True)
return np.asarray(action, dtype=np.float32).copy()
def _mean_or_zero(values: list[float]) -> float:
return float(np.mean(values)) if values else 0.0
def _stats_or_zero(values: list[float]) -> dict[str, float]:
if not values:
return {
'mean': 0.0,
'std': 0.0,
'min': 0.0,
'max': 0.0,
}
array = np.asarray(values, dtype=np.float64)
return {
'mean': float(array.mean()),
'std': float(array.std()),
'min': float(array.min()),
'max': float(array.max()),
}
def _summarize_timing_breakdown(
all_timings: dict[str, list[float]],
model_forward_flags: list[bool],
) -> dict[str, Any]:
model_forward_flags = [bool(flag) for flag in model_forward_flags]
return {
'count': int(len(model_forward_flags)),
'model_forward_count': int(sum(model_forward_flags)),
'all_steps_ms': {
stage: _stats_or_zero(values)
for stage, values in all_timings.items()
},
'model_forward_steps_ms': {
stage: _stats_or_zero(
[value for value, should_keep in zip(values, model_forward_flags) if should_keep]
)
for stage, values in all_timings.items()
},
}
_TIMING_SAMPLE_KEYS = (
'obs_read_time_ms',
'preprocess_time_ms',
'inference_time_ms',
'env_step_time_ms',
'total_time_ms',
)
def _empty_merge_state() -> dict[str, list[float] | list[bool]]:
return {
**{key: [] for key in _TIMING_SAMPLE_KEYS},
'model_forward_flags': [],
}
def _normalize_num_workers(num_workers: int, num_episodes: int) -> int:
num_episodes = max(int(num_episodes), 0)
if num_episodes == 0:
return 0
return min(max(int(num_workers), 1), num_episodes)
def _split_episode_indices(num_episodes: int, num_workers: int) -> list[list[int]]:
active_workers = _normalize_num_workers(num_workers=num_workers, num_episodes=num_episodes)
if active_workers == 0:
return []
episode_indices = np.arange(int(num_episodes), dtype=np.int32)
return [
chunk.tolist()
for chunk in np.array_split(episode_indices, active_workers)
if len(chunk) > 0
]
def _plan_episode_box_poses(
num_episodes: int,
sampler=None,
) -> list[np.ndarray]:
if sampler is None:
sampler = sample_transfer_pose
return [
np.asarray(sampler(), dtype=np.float32).copy()
for _ in range(int(num_episodes))
]
def _merge_worker_summaries(
worker_summaries: list[dict[str, Any]],
artifact_paths: dict[str, Optional[str]],
) -> dict[str, Any]:
merged_episodes = []
merged_state = _empty_merge_state()
for worker_summary in worker_summaries:
merged_episodes.extend(worker_summary.get('episodes', []))
merge_state = worker_summary.get('_merge_state', {})
for key in _TIMING_SAMPLE_KEYS:
merged_state[key].extend(float(value) for value in merge_state.get(key, []))
merged_state['model_forward_flags'].extend(
bool(value) for value in merge_state.get('model_forward_flags', [])
)
merged_episodes = sorted(
merged_episodes,
key=lambda episode: int(episode.get('episode_index', 0)),
)
episode_rewards = [
float(episode.get('episode_reward', 0.0))
for episode in merged_episodes
]
episode_max_rewards = [
(
float(episode['episode_max_reward'])
if episode.get('episode_max_reward') is not None
else None
)
for episode in merged_episodes
]
valid_max_rewards = [
value for value in episode_max_rewards
if value is not None
]
summary = {
'num_episodes': len(merged_episodes),
'episode_rewards': episode_rewards,
'episode_max_rewards': episode_max_rewards,
'avg_reward': float(np.mean(episode_rewards)) if episode_rewards else 0.0,
'avg_max_reward': float(np.mean(valid_max_rewards)) if valid_max_rewards else 0.0,
'episodes': merged_episodes,
'artifact_dir': artifact_paths.get('output_dir') if artifact_paths else None,
'artifacts': artifact_paths,
}
if merged_episodes:
summary.update({
'avg_inference_fps': float(np.mean([
float(episode.get('inference_fps', 0.0))
for episode in merged_episodes
])),
'avg_control_fps': float(np.mean([
float(episode.get('control_fps', 0.0))
for episode in merged_episodes
])),
'avg_obs_read_time_ms': _mean_or_zero(merged_state['obs_read_time_ms']),
'avg_preprocess_time_ms': _mean_or_zero(merged_state['preprocess_time_ms']),
'avg_inference_time_ms': _mean_or_zero(merged_state['inference_time_ms']),
'avg_env_step_time_ms': _mean_or_zero(merged_state['env_step_time_ms']),
'avg_total_time_ms': _mean_or_zero(merged_state['total_time_ms']),
'timing_summary': _summarize_timing_breakdown(
{
'obs_read': merged_state['obs_read_time_ms'],
'preprocess': merged_state['preprocess_time_ms'],
'inference': merged_state['inference_time_ms'],
'env_step': merged_state['env_step_time_ms'],
'loop_total': merged_state['total_time_ms'],
},
merged_state['model_forward_flags'],
),
})
return summary
def _json_friendly(value: Any) -> Any:
if isinstance(value, dict):
return {str(key): _json_friendly(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_json_friendly(item) for item in value]
if isinstance(value, Path):
return str(value)
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, (np.integer, np.floating)):
return value.item()
return value
def _resolve_artifact_paths(
eval_cfg: DictConfig,
output_dir_override: Optional[str] = None,
) -> dict[str, Optional[str]]:
save_timing = bool(eval_cfg.get('save_timing', False))
save_trajectory = bool(
eval_cfg.get('save_trajectory', False) or eval_cfg.get('save_trajectory_npz', False)
)
wants_artifacts = any([
bool(eval_cfg.get('save_artifacts', False)),
save_timing,
save_trajectory,
bool(eval_cfg.get('record_video', False)),
])
output_dir: Optional[Path] = None
if wants_artifacts:
if output_dir_override:
output_dir = Path(str(output_dir_override)).expanduser().resolve()
else:
artifact_dir = eval_cfg.get('artifact_dir', None)
if artifact_dir:
output_dir = Path(str(artifact_dir)).expanduser().resolve()
else:
ckpt_stem = Path(str(eval_cfg.ckpt_path)).stem or 'rollout'
timestamp = time.strftime('%Y%m%d-%H%M%S')
output_dir = (Path.cwd() / 'rollout_artifacts' / f'{ckpt_stem}-{timestamp}').resolve()
output_dir.mkdir(parents=True, exist_ok=True)
video_camera_name = None
if bool(eval_cfg.get('record_video', False)):
configured_camera_name = eval_cfg.get('video_camera_name', None)
if configured_camera_name is None:
configured_camera_name = eval_cfg.get('video_camera', None)
if configured_camera_name is not None:
video_camera_name = str(configured_camera_name)
elif eval_cfg.get('camera_names'):
video_camera_name = str(eval_cfg.camera_names[0])
else:
raise ValueError('record_video=true requires eval.video_camera_name or a non-empty eval.camera_names')
return {
'output_dir': str(output_dir) if output_dir is not None else None,
'summary_json': (
str(output_dir / 'rollout_summary.json')
if output_dir is not None and bool(eval_cfg.get('save_summary_json', False))
else None
),
'timing_json': (
str(output_dir / 'timing.json')
if output_dir is not None and save_timing
else None
),
'trajectory_npz': (
str(output_dir / 'trajectory.npz')
if output_dir is not None and save_trajectory
else None
),
'video_mp4': (
str(output_dir / f'rollout_{video_camera_name}.mp4')
if output_dir is not None and bool(eval_cfg.get('record_video', False))
and video_camera_name is not None
else None
),
'video_camera_name': video_camera_name,
}
def _get_video_frame(obs: Dict, camera_name: Optional[str]) -> Optional[np.ndarray]:
if camera_name is None:
return None
frame = obs['images'][camera_name]
frame = np.asarray(frame)
if frame.ndim != 3 or frame.shape[2] != 3:
raise ValueError(
f'Video frame for camera {camera_name} must have shape (H, W, 3), got {frame.shape}'
)
if frame.dtype != np.uint8:
frame = np.clip(frame, 0, 255).astype(np.uint8)
return frame
def _open_video_writer(output_path: str, frame_size: tuple[int, int], fps: int):
import cv2
output_path = str(output_path)
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
writer = cv2.VideoWriter(output_path, fourcc, float(fps), frame_size)
if not writer.isOpened():
raise RuntimeError(f'无法打开视频输出: {output_path}')
return writer
class _RolloutVideoRecorder:
def __init__(self, output_path: Optional[str], fps: int):
self.output_path = output_path
self.fps = int(fps)
self.writer = None
def write(self, frame: Optional[np.ndarray]):
if self.output_path is None or frame is None:
return
if self.writer is None:
frame_size = (int(frame.shape[1]), int(frame.shape[0]))
self.writer = _open_video_writer(self.output_path, frame_size, self.fps)
self.writer.write(frame)
def close(self):
if self.writer is not None:
self.writer.release()
self.writer = None
def _read_body_pose(env, body_name: str):
try:
if callable(getattr(env, 'getBodyPos', None)) and callable(getattr(env, 'getBodyQuat', None)):
pos = env.getBodyPos(body_name)
quat = env.getBodyQuat(body_name)
else:
body = env.mj_data.body(body_name)
pos = body.xpos
quat = body.xquat
except Exception:
return None
return {
'pos': np.asarray(pos, dtype=np.float32).copy(),
'quat': np.asarray(quat, dtype=np.float32).copy(),
}
def _get_executed_ee_poses(env) -> dict[str, np.ndarray]:
candidates = {
'left_link7': ('left_link7', 'eef_left'),
'right_link7': ('right_link7', 'eef_right'),
'eef_left': ('eef_left', 'left_link7'),
'eef_right': ('eef_right', 'right_link7'),
}
poses = {}
for body_key, body_names in candidates.items():
pose = None
for body_name in body_names:
pose = _read_body_pose(env, body_name)
if pose is not None:
break
if pose is None:
pose = {
'pos': np.full(3, np.nan, dtype=np.float32),
'quat': np.full(4, np.nan, dtype=np.float32),
}
poses[f'{body_key}_pos'] = pose['pos']
poses[f'{body_key}_quat'] = pose['quat']
return poses
def _empty_rollout_trajectory() -> dict[str, list]:
return {
'episode_index': [],
'step': [],
'reward': [],
'raw_action': [],
'applied_action': [],
'executed_left_link7_pos': [],
'executed_left_link7_quat': [],
'executed_right_link7_pos': [],
'executed_right_link7_quat': [],
'executed_eef_left_pos': [],
'executed_eef_left_quat': [],
'executed_eef_right_pos': [],
'executed_eef_right_quat': [],
'model_inference_triggered': [],
'obs_read_time_ms': [],
'preprocess_time_ms': [],
'inference_time_ms': [],
'env_step_time_ms': [],
'total_time_ms': [],
}
def _append_rollout_step(
storage: dict[str, list],
episode_index: int,
timestep: int,
reward: Optional[float],
raw_action: np.ndarray,
executed_action: np.ndarray,
executed_poses: dict[str, np.ndarray],
timing_ms: dict[str, float],
model_inference_triggered: bool,
):
storage['episode_index'].append(int(episode_index))
storage['step'].append(int(timestep))
storage['reward'].append(float(reward) if reward is not None else np.nan)
storage['raw_action'].append(raw_action.astype(np.float32, copy=True))
storage['applied_action'].append(executed_action.astype(np.float32, copy=True))
storage['executed_left_link7_pos'].append(executed_poses['left_link7_pos'])
storage['executed_left_link7_quat'].append(executed_poses['left_link7_quat'])
storage['executed_right_link7_pos'].append(executed_poses['right_link7_pos'])
storage['executed_right_link7_quat'].append(executed_poses['right_link7_quat'])
storage['executed_eef_left_pos'].append(executed_poses['eef_left_pos'])
storage['executed_eef_left_quat'].append(executed_poses['eef_left_quat'])
storage['executed_eef_right_pos'].append(executed_poses['eef_right_pos'])
storage['executed_eef_right_quat'].append(executed_poses['eef_right_quat'])
storage['model_inference_triggered'].append(bool(model_inference_triggered))
for key, value in timing_ms.items():
storage[key].append(float(value))
def _save_rollout_trajectory_npz(output_path: str, storage: dict[str, list]):
step = np.asarray(storage['step'], dtype=np.int32)
raw_action = np.asarray(storage['raw_action'], dtype=np.float32)
applied_action = np.asarray(storage['applied_action'], dtype=np.float32)
executed_left_link7_pos = np.asarray(storage['executed_left_link7_pos'], dtype=np.float32)
executed_left_link7_quat = np.asarray(storage['executed_left_link7_quat'], dtype=np.float32)
executed_right_link7_pos = np.asarray(storage['executed_right_link7_pos'], dtype=np.float32)
executed_right_link7_quat = np.asarray(storage['executed_right_link7_quat'], dtype=np.float32)
executed_eef_left_pos = np.asarray(storage['executed_eef_left_pos'], dtype=np.float32)
executed_eef_left_quat = np.asarray(storage['executed_eef_left_quat'], dtype=np.float32)
executed_eef_right_pos = np.asarray(storage['executed_eef_right_pos'], dtype=np.float32)
executed_eef_right_quat = np.asarray(storage['executed_eef_right_quat'], dtype=np.float32)
np.savez_compressed(
output_path,
episode_index=np.asarray(storage['episode_index'], dtype=np.int32),
step=step,
timestep=step,
reward=np.asarray(storage['reward'], dtype=np.float32),
raw_action=raw_action,
raw_predicted_ee_action=raw_action,
applied_action=applied_action,
executed_ee_action=applied_action,
executed_left_link7_pos=executed_left_link7_pos,
executed_left_link7_quat=executed_left_link7_quat,
executed_right_link7_pos=executed_right_link7_pos,
executed_right_link7_quat=executed_right_link7_quat,
executed_eef_left_pos=executed_eef_left_pos,
executed_eef_left_quat=executed_eef_left_quat,
executed_eef_right_pos=executed_eef_right_pos,
executed_eef_right_quat=executed_eef_right_quat,
left_ee_pos=executed_eef_left_pos,
left_ee_quat=executed_eef_left_quat,
right_ee_pos=executed_eef_right_pos,
right_ee_quat=executed_eef_right_quat,
model_inference_triggered=np.asarray(storage['model_inference_triggered'], dtype=bool),
obs_read_time_ms=np.asarray(storage['obs_read_time_ms'], dtype=np.float32),
preprocess_time_ms=np.asarray(storage['preprocess_time_ms'], dtype=np.float32),
inference_time_ms=np.asarray(storage['inference_time_ms'], dtype=np.float32),
env_step_time_ms=np.asarray(storage['env_step_time_ms'], dtype=np.float32),
total_time_ms=np.asarray(storage['total_time_ms'], dtype=np.float32),
)
def _save_summary_json(output_path: str, summary: dict[str, Any]):
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(_json_friendly(summary), f, ensure_ascii=False, indent=2)
def _public_summary(summary: dict[str, Any]) -> dict[str, Any]:
public_summary = dict(summary)
public_summary.pop('_merge_state', None)
return _json_friendly(public_summary)
class ActionSmoother:
"""
动作平滑器(指数移动平均)
用于平滑执行动作以获得更稳定的控制
"""
def __init__(self, alpha: float = 0.3):
"""
Args:
alpha: 平滑系数 (0-1),值越大越重视当前动作
"""
self.alpha = alpha
self.prev_action = None
def smooth(self, action: np.ndarray) -> np.ndarray:
"""
平滑动作
Args:
action: 当前动作
Returns:
平滑后的动作
"""
if self.prev_action is None:
smoothed = action
else:
smoothed = self.alpha * action + (1 - self.alpha) * self.prev_action
self.prev_action = smoothed
return smoothed
def reset(self):
"""重置平滑器状态"""
self.prev_action = None
def _close_env(env):
if env is None:
return
if hasattr(env, 'exit_flag'):
env.exit_flag = True
cam_thread = getattr(env, 'cam_thread', None)
if cam_thread is not None and hasattr(cam_thread, 'join'):
cam_thread.join(timeout=1.0)
viewer = getattr(env, 'viewer', None)
if viewer is not None and hasattr(viewer, 'close'):
viewer.close()
def _print_eval_config(cfg: DictConfig):
# 打印配置
print("=" * 80)
print("VLA 评估配置:")
print("=" * 80)
print(OmegaConf.to_yaml(cfg))
print("=" * 80)
def _build_episode_plans(
num_episodes: int,
box_poses: Optional[list[np.ndarray]] = None,
) -> list[dict[str, Any]]:
if box_poses is None:
return [
{
'episode_index': int(episode_index),
}
for episode_index in range(int(num_episodes))
]
return [
{
'episode_index': int(episode_index),
'box_pos': np.asarray(box_pos, dtype=np.float32).copy(),
}
for episode_index, box_pos in enumerate(box_poses)
]
def _run_eval_episode_plans(
cfg: DictConfig,
episode_plans: list[dict[str, Any]],
policy_runner,
worker_index: int = 0,
artifact_paths: Optional[dict[str, Optional[str]]] = None,
show_progress: bool = True,
) -> dict[str, Any]:
eval_cfg = cfg.eval
device = str(eval_cfg.device)
camera_names = list(eval_cfg.camera_names)
artifact_paths = artifact_paths or _resolve_artifact_paths(eval_cfg)
video_recorder = _RolloutVideoRecorder(
output_path=artifact_paths['video_mp4'],
fps=int(eval_cfg.get('video_fps', 30)),
)
rollout_trajectory = _empty_rollout_trajectory()
merge_state = _empty_merge_state()
# 可选:动作平滑器
smoother = ActionSmoother(alpha=eval_cfg.smooth_alpha) if eval_cfg.use_smoothing else None
# =========================================================================
# 创建环境
# =========================================================================
env = None
try:
env = make_sim_env(eval_cfg.task_name, headless=eval_cfg.headless)
# =========================================================================
# 运行评估回合
# =========================================================================
all_stats = []
episode_rewards = []
episode_max_rewards = []
for plan in episode_plans:
episode_idx = int(plan['episode_index'])
box_pos = plan.get('box_pos')
if box_pos is None:
box_pos = sample_transfer_pose()
box_pos = np.asarray(box_pos, dtype=np.float32)
if show_progress:
print(f"\n{'='*60}")
print(f"回合 {episode_idx + 1}/{eval_cfg.num_episodes}")
print(f"{'='*60}\n")
env.reset(box_pos)
# 为新回合重置 rollout policy 状态
policy_runner.reset()
if smoother:
smoother.reset()
# 计时统计
obs_read_times_ms = []
preprocess_times_ms = []
inference_times_ms = []
env_step_times_ms = []
total_times_ms = []
model_forward_flags = []
episode_reward = 0.0
episode_max_reward = float('-inf')
with torch.inference_mode():
episode_iterator = range(eval_cfg.max_timesteps)
if show_progress:
episode_iterator = tqdm(episode_iterator, desc=f"回合 {episode_idx + 1}")
for t in episode_iterator:
start_total = time.perf_counter()
# 从环境获取观测
obs = env._get_image_obs()
qpos_obs = env._get_qpos_obs()
obs['qpos'] = qpos_obs['qpos']
end_obs_read = time.perf_counter()
video_frame = _get_video_frame(obs, artifact_paths['video_camera_name'])
video_recorder.write(video_frame)
# 准备给 agent 的观测
observation = prepare_observation(obs, camera_names)
end_preprocess = time.perf_counter()
# 选择动作(本地 agent 或远端 inference server
start_inference = time.perf_counter()
action, model_inference_triggered = policy_runner.select_action(
observation,
episode_index=episode_idx,
timestep=t,
)
if (
getattr(policy_runner, 'uses_local_model', False)
and _is_cuda_device(device)
and torch.cuda.is_available()
):
torch.cuda.synchronize()
end_inference = time.perf_counter()
# 转换为 numpy
raw_action = _to_numpy_action(action)
# 调试:打印当前时间步的动作(由配置控制)
if eval_cfg.get('verbose_action', False):
print(f"\n[Step {t:3d}] 预测动作: {raw_action}")
print(f" - 动作形状: {raw_action.shape}")
print(f" - 动作范围: [{raw_action.min():.4f}, {raw_action.max():.4f}]")
print(f" - 动作均值: {raw_action.mean():.4f}, 标准差: {raw_action.std():.4f}")
# 可选:平滑动作
executed_action = raw_action.copy()
if smoother:
executed_action = smoother.smooth(executed_action)
# 执行动作
start_env_step = time.perf_counter()
execute_policy_action(env, executed_action)
end_env_step = time.perf_counter()
executed_poses = _get_executed_ee_poses(env)
reward = getattr(env, 'rew', None)
if reward is not None:
reward = float(reward)
episode_reward += reward
episode_max_reward = max(episode_max_reward, reward)
if not eval_cfg.headless:
env.render()
end_total = time.perf_counter()
step_timing_ms = {
'obs_read_time_ms': (end_obs_read - start_total) * 1000.0,
'preprocess_time_ms': (end_preprocess - end_obs_read) * 1000.0,
'inference_time_ms': (end_inference - start_inference) * 1000.0,
'env_step_time_ms': (end_env_step - start_env_step) * 1000.0,
'total_time_ms': (end_total - start_total) * 1000.0,
}
# 记录计时
obs_read_times_ms.append(step_timing_ms['obs_read_time_ms'])
preprocess_times_ms.append(step_timing_ms['preprocess_time_ms'])
inference_times_ms.append(step_timing_ms['inference_time_ms'])
env_step_times_ms.append(step_timing_ms['env_step_time_ms'])
total_times_ms.append(step_timing_ms['total_time_ms'])
model_forward_flags.append(bool(model_inference_triggered))
merge_state['obs_read_time_ms'].append(step_timing_ms['obs_read_time_ms'])
merge_state['preprocess_time_ms'].append(step_timing_ms['preprocess_time_ms'])
merge_state['inference_time_ms'].append(step_timing_ms['inference_time_ms'])
merge_state['env_step_time_ms'].append(step_timing_ms['env_step_time_ms'])
merge_state['total_time_ms'].append(step_timing_ms['total_time_ms'])
merge_state['model_forward_flags'].append(bool(model_inference_triggered))
if artifact_paths['trajectory_npz'] is not None:
_append_rollout_step(
rollout_trajectory,
episode_index=episode_idx,
timestep=t,
reward=reward,
raw_action=raw_action,
executed_action=executed_action,
executed_poses=executed_poses,
timing_ms=step_timing_ms,
model_inference_triggered=model_inference_triggered,
)
# =========================================================================
# 打印回合统计
# =========================================================================
avg_obs_read_time_ms = _mean_or_zero(obs_read_times_ms)
avg_preprocess_time_ms = _mean_or_zero(preprocess_times_ms)
avg_inference_time_ms = _mean_or_zero(inference_times_ms)
avg_env_step_time_ms = _mean_or_zero(env_step_times_ms)
avg_total_time_ms = _mean_or_zero(total_times_ms)
timing_breakdown = _summarize_timing_breakdown(
{
'obs_read': obs_read_times_ms,
'preprocess': preprocess_times_ms,
'inference': inference_times_ms,
'env_step': env_step_times_ms,
'loop_total': total_times_ms,
},
model_forward_flags,
)
episode_artifact_paths = {
'video': artifact_paths['video_mp4'],
'trajectory': artifact_paths['trajectory_npz'],
'timing': artifact_paths['timing_json'] or artifact_paths['summary_json'],
}
stats = {
'worker_index': int(worker_index),
'episode_index': int(episode_idx),
'inference_fps': 1000.0 / avg_inference_time_ms if avg_inference_time_ms > 0 else 0.0,
'control_fps': 1000.0 / avg_total_time_ms if avg_total_time_ms > 0 else 0.0,
'avg_obs_read_time_ms': avg_obs_read_time_ms,
'avg_preprocess_time_ms': avg_preprocess_time_ms,
'avg_inference_time_ms': avg_inference_time_ms,
'avg_env_step_time_ms': avg_env_step_time_ms,
'avg_total_time_ms': avg_total_time_ms,
'num_inferences': int(sum(model_forward_flags)),
'num_model_forwards': int(sum(model_forward_flags)),
'num_steps': len(total_times_ms),
'episode_reward': float(episode_reward),
'episode_max_reward': (
float(episode_max_reward) if episode_max_reward != float('-inf') else None
),
'artifact_paths': episode_artifact_paths,
'timing_breakdown_ms': timing_breakdown['all_steps_ms'],
'timing_summary': timing_breakdown,
}
all_stats.append(stats)
episode_rewards.append(float(episode_reward))
episode_max_rewards.append(
float(episode_max_reward) if episode_max_reward != float('-inf') else None
)
if show_progress:
print(f"\n回合 {episode_idx + 1} 完成 ({eval_cfg.max_timesteps} 时间步)")
print(f" 模型推理 FPS: {stats['inference_fps']:.2f} Hz")
print(f" 控制循环 FPS: {stats['control_fps']:.2f} Hz")
print(f" 平均读观测时间: {stats['avg_obs_read_time_ms']:.2f} ms")
print(f" 平均预处理时间: {stats['avg_preprocess_time_ms']:.2f} ms")
print(f" 平均推理时间: {stats['avg_inference_time_ms']:.2f} ms")
print(f" 平均环境步进时间: {stats['avg_env_step_time_ms']:.2f} ms")
print(f" 平均总时间: {stats['avg_total_time_ms']:.2f} ms")
print(f" 总推理次数: {stats['num_inferences']}")
print(f" 回合累计奖励: {stats['episode_reward']:.2f}")
# =========================================================================
# 总体统计
# =========================================================================
if show_progress:
print(f"\n{'='*60}")
print("评估完成!")
print(f"{'='*60}")
valid_max_rewards = [
reward for reward in episode_max_rewards
if reward is not None
]
summary = {
'num_episodes': len(episode_plans),
'episode_rewards': episode_rewards,
'episode_max_rewards': episode_max_rewards,
'avg_reward': float(np.mean(episode_rewards)) if episode_rewards else 0.0,
'avg_max_reward': float(np.mean(valid_max_rewards)) if valid_max_rewards else 0.0,
'episodes': all_stats,
'artifact_dir': artifact_paths['output_dir'],
'artifacts': artifact_paths,
'_merge_state': merge_state,
}
if all_stats:
avg_inference_fps = np.mean([s['inference_fps'] for s in all_stats])
avg_control_fps = np.mean([s['control_fps'] for s in all_stats])
avg_obs_read_time = _mean_or_zero(merge_state['obs_read_time_ms'])
avg_preprocess_time = _mean_or_zero(merge_state['preprocess_time_ms'])
avg_inference_time = _mean_or_zero(merge_state['inference_time_ms'])
avg_env_step_time = _mean_or_zero(merge_state['env_step_time_ms'])
avg_total_time = _mean_or_zero(merge_state['total_time_ms'])
summary.update({
'avg_inference_fps': float(avg_inference_fps),
'avg_control_fps': float(avg_control_fps),
'avg_obs_read_time_ms': float(avg_obs_read_time),
'avg_preprocess_time_ms': float(avg_preprocess_time),
'avg_inference_time_ms': float(avg_inference_time),
'avg_env_step_time_ms': float(avg_env_step_time),
'avg_total_time_ms': float(avg_total_time),
'timing_summary': _summarize_timing_breakdown(
{
'obs_read': merge_state['obs_read_time_ms'],
'preprocess': merge_state['preprocess_time_ms'],
'inference': merge_state['inference_time_ms'],
'env_step': merge_state['env_step_time_ms'],
'loop_total': merge_state['total_time_ms'],
},
merge_state['model_forward_flags'],
),
})
if show_progress:
print(f"\n总体统计 ({len(episode_plans)} 个回合):")
print(f" 平均模型推理 FPS: {avg_inference_fps:.2f} Hz")
print(f" 平均控制循环 FPS: {avg_control_fps:.2f} Hz")
print(f" 平均读观测时间: {avg_obs_read_time:.2f} ms")
print(f" 平均预处理时间: {avg_preprocess_time:.2f} ms")
print(f" 平均推理时间: {avg_inference_time:.2f} ms")
print(f" 平均环境步进时间: {avg_env_step_time:.2f} ms")
print(f" 平均总时间: {avg_total_time:.2f} ms")
print(f" 平均累计奖励: {summary['avg_reward']:.2f}")
if artifact_paths['trajectory_npz'] is not None:
_save_rollout_trajectory_npz(artifact_paths['trajectory_npz'], rollout_trajectory)
public_summary = _public_summary(summary)
if artifact_paths['summary_json'] is not None:
_save_summary_json(artifact_paths['summary_json'], public_summary)
if artifact_paths['timing_json'] is not None:
_save_summary_json(artifact_paths['timing_json'], public_summary.get('timing_summary', {}))
if show_progress:
print()
return summary
finally:
video_recorder.close()
_close_env(env)
def _run_eval_worker(
cfg: DictConfig,
episode_plans: list[dict[str, Any]],
worker_index: int = 0,
artifact_paths: Optional[dict[str, Optional[str]]] = None,
show_progress: bool = True,
) -> dict[str, Any]:
eval_cfg = cfg.eval
log.info(f"🚀 从 {eval_cfg.ckpt_path} 加载模型...")
agent, _dataset_stats = load_checkpoint(
ckpt_path=eval_cfg.ckpt_path,
agent_cfg=cfg.agent,
device=str(eval_cfg.device),
)
policy_runner = _LocalPolicyRunner(agent)
return _run_eval_episode_plans(
cfg,
episode_plans=episode_plans,
policy_runner=policy_runner,
worker_index=worker_index,
artifact_paths=artifact_paths,
show_progress=show_progress,
)
def _run_remote_eval_worker(
cfg: DictConfig,
episode_plans: list[dict[str, Any]],
*,
worker_index: int,
server_index: int,
request_queue,
response_queue,
artifact_paths: Optional[dict[str, Optional[str]]] = None,
show_progress: bool = False,
) -> dict[str, Any]:
eval_cfg = cfg.eval
agent_cfg = cfg.agent
num_action_steps = int(agent_cfg.get('num_action_steps', eval_cfg.get('num_queries', 1)))
policy_runner = _RemotePolicyRunner(
worker_index=worker_index,
server_index=server_index,
request_queue=request_queue,
response_queue=response_queue,
camera_names=_resolve_policy_camera_names(cfg),
obs_horizon=int(agent_cfg.get('obs_horizon', eval_cfg.obs_horizon)),
num_action_steps=num_action_steps,
response_timeout_s=float(eval_cfg.get('response_timeout_s', 300.0)),
)
return _run_eval_episode_plans(
cfg,
episode_plans=episode_plans,
policy_runner=policy_runner,
worker_index=worker_index,
artifact_paths=artifact_paths,
show_progress=show_progress,
)
def _run_eval_serial(cfg: DictConfig):
eval_cfg = cfg.eval
artifact_paths = _resolve_artifact_paths(eval_cfg)
episode_plans = _build_episode_plans(eval_cfg.num_episodes)
summary = _run_eval_worker(
cfg,
episode_plans=episode_plans,
worker_index=0,
artifact_paths=artifact_paths,
show_progress=True,
)
return _public_summary(summary)
def _validate_parallel_eval_cfg(eval_cfg: DictConfig):
if not bool(eval_cfg.get('headless', False)):
raise ValueError('eval.num_workers > 1 requires eval.headless=true')
unsupported_exports = [
flag_name
for flag_name in (
'record_video',
'save_trajectory',
'save_trajectory_npz',
'save_trajectory_image',
)
if bool(eval_cfg.get(flag_name, False))
]
if unsupported_exports:
joined_flags = ', '.join(unsupported_exports)
raise ValueError(
'eval.num_workers > 1 does not yet support parallel export for '
f'{joined_flags}'
)
def _is_cuda_device(device: Any) -> bool:
return str(device).lower().startswith('cuda')
def _resolve_cuda_devices(eval_cfg: DictConfig) -> list[int]:
if not _is_cuda_device(eval_cfg.get('device', 'cpu')):
return []
configured_devices = eval_cfg.get('cuda_devices', None)
if configured_devices is None:
return [0]
resolved_devices = [int(device_index) for device_index in configured_devices]
if not resolved_devices:
raise ValueError('eval.cuda_devices must not be empty when eval.device is CUDA')
if any(device_index < 0 for device_index in resolved_devices):
raise ValueError('eval.cuda_devices must contain non-negative logical CUDA device indices')
return resolved_devices
def _run_spawn_jobs(
payloads: list[dict[str, Any]],
max_workers: int,
worker_fn,
) -> list[Any]:
if not payloads:
return []
ctx = multiprocessing.get_context('spawn')
results = []
with concurrent.futures.ProcessPoolExecutor(
max_workers=int(max_workers),
mp_context=ctx,
) as executor:
future_to_payload = {
executor.submit(worker_fn, payload): payload
for payload in payloads
}
try:
for future in concurrent.futures.as_completed(future_to_payload):
results.append(future.result())
except Exception:
for future in future_to_payload:
future.cancel()
raise
return results
def _run_eval_worker_entry(payload: dict[str, Any]) -> dict[str, Any]:
if payload.get('_spawn_probe', False):
return {
'probe_value': int(payload['probe_value']),
'worker_index': int(payload.get('worker_index', -1)),
}
cfg = OmegaConf.create(payload['cfg'])
artifact_paths = _resolve_artifact_paths(
cfg.eval,
output_dir_override=payload.get('artifact_dir'),
)
return _run_eval_worker(
cfg,
episode_plans=list(payload.get('episode_plans', [])),
worker_index=int(payload.get('worker_index', 0)),
artifact_paths=artifact_paths,
show_progress=False,
)
def _build_parallel_worker_payloads(
cfg: DictConfig,
artifact_paths: dict[str, Optional[str]],
) -> tuple[list[dict[str, Any]], int]:
eval_cfg = cfg.eval
requested_workers = int(eval_cfg.get('num_workers', 1))
episode_splits = _split_episode_indices(
num_episodes=int(eval_cfg.num_episodes),
num_workers=requested_workers,
)
box_poses = _plan_episode_box_poses(int(eval_cfg.num_episodes))
resolved_cfg = OmegaConf.to_container(cfg, resolve=True)
payloads = []
workers_dir = None
if artifact_paths.get('output_dir') is not None:
workers_dir = Path(str(artifact_paths['output_dir'])) / 'workers'
workers_dir.mkdir(parents=True, exist_ok=True)
artifact_paths['workers_dir'] = str(workers_dir)
for worker_index, episode_indices in enumerate(episode_splits):
worker_artifact_dir = None
if workers_dir is not None:
worker_artifact_dir = workers_dir / f'worker_{worker_index:02d}'
worker_artifact_dir.mkdir(parents=True, exist_ok=True)
worker_cfg = json.loads(json.dumps(resolved_cfg))
worker_cfg['eval']['artifact_dir'] = (
str(worker_artifact_dir) if worker_artifact_dir is not None else None
)
payloads.append({
'cfg': worker_cfg,
'worker_index': int(worker_index),
'artifact_dir': str(worker_artifact_dir) if worker_artifact_dir is not None else None,
'episode_plans': [
{
'episode_index': int(episode_index),
'box_pos': np.asarray(box_poses[episode_index], dtype=np.float32).tolist(),
}
for episode_index in episode_indices
],
})
return payloads, len(episode_splits)
def _build_cuda_server_payloads(
cfg: DictConfig,
worker_payloads: list[dict[str, Any]],
cuda_devices: list[int],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
resolved_cfg = OmegaConf.to_container(cfg, resolve=True)
server_payloads = [
{
'cfg': json.loads(json.dumps(resolved_cfg)),
'server_index': int(server_index),
'device_index': int(device_index),
'worker_indices': [],
}
for server_index, device_index in enumerate(cuda_devices)
]
assigned_workers = []
for worker_payload in worker_payloads:
server_index = int(worker_payload['worker_index']) % len(server_payloads)
assigned_payload = dict(worker_payload)
assigned_payload['server_index'] = server_index
assigned_payload['device_index'] = int(server_payloads[server_index]['device_index'])
assigned_workers.append(assigned_payload)
server_payloads[server_index]['worker_indices'].append(int(worker_payload['worker_index']))
return server_payloads, assigned_workers
def _inference_server_main(payload: dict[str, Any]) -> None:
request_queue = payload['request_queue']
response_queues = payload['response_queues']
server_index = int(payload.get('server_index', -1))
if payload.get('_spawn_probe', False):
while True:
message = request_queue.get()
if message.get('type') == 'shutdown_server':
return
if message.get('type') != 'predict_chunk':
continue
response_queues[int(message['worker_index'])].put({
'type': 'predict_chunk_result',
'server_index': server_index,
'actions': np.asarray([[[11.0], [22.0], [33.0]]], dtype=np.float32),
})
return
result_queue = payload.get('result_queue')
cfg = OmegaConf.create(payload['cfg'])
device = f'cuda:{int(payload["device_index"])}'
try:
agent, _dataset_stats = load_checkpoint(
ckpt_path=cfg.eval.ckpt_path,
agent_cfg=cfg.agent,
device=device,
)
if result_queue is not None:
result_queue.put({
'kind': 'server_ready',
'server_index': server_index,
})
while True:
message = request_queue.get()
message_type = message.get('type')
if message_type == 'shutdown_server':
return
if message_type != 'predict_chunk':
raise ValueError(f'unknown server message type: {message_type}')
batch = _deserialize_policy_batch(message['batch'], device=device)
with torch.inference_mode():
actions = agent.predict_action_chunk(batch)
if torch.cuda.is_available():
torch.cuda.synchronize()
response_queues[int(message['worker_index'])].put({
'type': 'predict_chunk_result',
'server_index': server_index,
'actions': actions.detach().cpu().numpy().astype(np.float32, copy=True),
})
except Exception as exc:
if result_queue is not None:
result_queue.put({
'kind': 'server_error',
'server_index': server_index,
'error': str(exc),
})
raise
def _env_worker_main(payload: dict[str, Any]) -> None:
worker_index = int(payload.get('worker_index', -1))
server_index = int(payload.get('server_index', -1))
request_queue = payload['request_queue']
response_queue = payload['response_queue']
result_queue = payload.get('result_queue')
if payload.get('_spawn_probe', False):
request_queue.put({
'type': 'predict_chunk',
'worker_index': worker_index,
'server_index': server_index,
'episode_index': 0,
'timestep': 0,
'batch': {
'qpos': np.zeros((1, 1, 1), dtype=np.float32),
'images': {},
},
})
response = response_queue.get(timeout=10.0)
result_queue.put({
'kind': 'worker_result',
'worker_index': worker_index,
'summary': {
'probe_worker_index': worker_index,
'probe_server_index': server_index,
'probe_actions': np.asarray(response['actions']).tolist(),
},
})
return
cfg = OmegaConf.create(payload['cfg'])
artifact_paths = _resolve_artifact_paths(
cfg.eval,
output_dir_override=payload.get('artifact_dir'),
)
try:
summary = _run_remote_eval_worker(
cfg,
episode_plans=list(payload.get('episode_plans', [])),
worker_index=worker_index,
server_index=server_index,
request_queue=request_queue,
response_queue=response_queue,
artifact_paths=artifact_paths,
show_progress=False,
)
except Exception as exc:
if result_queue is not None:
result_queue.put({
'kind': 'worker_error',
'worker_index': worker_index,
'server_index': server_index,
'error': str(exc),
})
raise
if result_queue is not None:
result_queue.put({
'kind': 'worker_result',
'worker_index': worker_index,
'server_index': server_index,
'summary': summary,
})
def _shutdown_processes(processes: list[tuple[Any, dict[str, Any]]], *, terminate: bool = False) -> None:
for process, _payload in processes:
if terminate and process.is_alive():
process.terminate()
for process, _payload in processes:
process.join(timeout=5.0)
if process.is_alive():
process.terminate()
process.join(timeout=5.0)
def _run_cuda_parallel_processes(
server_payloads: list[dict[str, Any]],
worker_payloads: list[dict[str, Any]],
) -> list[dict[str, Any]]:
if not worker_payloads:
return []
ctx = multiprocessing.get_context('spawn')
result_queue = ctx.Queue()
response_queues = [ctx.Queue() for _ in range(len(worker_payloads))]
request_queues = {}
server_processes: list[tuple[Any, dict[str, Any]]] = []
worker_processes: list[tuple[Any, dict[str, Any]]] = []
should_terminate = False
startup_timeout_s = float(
server_payloads[0]['cfg']['eval'].get('server_startup_timeout_s', 300.0)
) if server_payloads else 300.0
try:
for server_payload in server_payloads:
request_queue = ctx.Queue()
request_queues[int(server_payload['server_index'])] = request_queue
process_payload = dict(server_payload)
process_payload['request_queue'] = request_queue
process_payload['response_queues'] = response_queues
process_payload['result_queue'] = result_queue
process = ctx.Process(
target=_inference_server_main,
args=(process_payload,),
name=f'eval-inference-server-{int(server_payload["server_index"]):02d}',
)
process.start()
server_processes.append((process, server_payload))
pending_servers = {int(payload['server_index']) for _process, payload in server_processes}
startup_deadline = time.monotonic() + startup_timeout_s
while pending_servers:
remaining = startup_deadline - time.monotonic()
if remaining <= 0:
raise RuntimeError(
'Timed out waiting for CUDA inference servers to become ready'
)
try:
message = result_queue.get(timeout=min(0.2, remaining))
except queue.Empty:
for process, payload in server_processes:
if process.exitcode not in (None, 0):
raise RuntimeError(
f'CUDA inference server {int(payload["server_index"])} exited with code '
f'{process.exitcode}'
)
continue
message_kind = message.get('kind')
if message_kind == 'server_ready':
pending_servers.discard(int(message['server_index']))
continue
if message_kind == 'server_error':
raise RuntimeError(
f'CUDA inference server {int(message["server_index"])} failed: {message.get("error")}'
)
raise RuntimeError(f'Unexpected CUDA startup message: {message_kind}')
for worker_payload in worker_payloads:
worker_index = int(worker_payload['worker_index'])
process_payload = dict(worker_payload)
process_payload['request_queue'] = request_queues[int(worker_payload['server_index'])]
process_payload['response_queue'] = response_queues[worker_index]
process_payload['result_queue'] = result_queue
process = ctx.Process(
target=_env_worker_main,
args=(process_payload,),
name=f'eval-env-worker-{worker_index:02d}',
)
process.start()
worker_processes.append((process, worker_payload))
pending_workers = {int(payload['worker_index']) for payload in worker_payloads}
worker_summaries = {}
while pending_workers:
try:
message = result_queue.get(timeout=0.2)
except queue.Empty:
for process, payload in server_processes:
if process.exitcode not in (None, 0):
raise RuntimeError(
f'CUDA inference server {int(payload["server_index"])} exited with code '
f'{process.exitcode}'
)
for process, payload in worker_processes:
worker_index = int(payload['worker_index'])
if worker_index in pending_workers and process.exitcode not in (None, 0):
raise RuntimeError(
f'CUDA rollout worker {worker_index} exited with code {process.exitcode}'
)
continue
message_kind = message.get('kind')
if message_kind == 'worker_result':
worker_index = int(message['worker_index'])
worker_summaries[worker_index] = message['summary']
pending_workers.discard(worker_index)
continue
if message_kind == 'worker_error':
raise RuntimeError(
f'CUDA rollout worker {int(message["worker_index"])} failed: {message.get("error")}'
)
if message_kind == 'server_error':
raise RuntimeError(
f'CUDA inference server {int(message["server_index"])} failed: {message.get("error")}'
)
raise RuntimeError(f'Unexpected CUDA parallel message: {message_kind}')
return [
worker_summaries[worker_index]
for worker_index in sorted(worker_summaries)
]
except Exception:
should_terminate = True
raise
finally:
for request_queue in request_queues.values():
try:
request_queue.put({'type': 'shutdown_server'})
except Exception:
continue
_shutdown_processes(worker_processes, terminate=should_terminate)
_shutdown_processes(server_processes, terminate=should_terminate)
def _run_eval_parallel(cfg: DictConfig):
eval_cfg = cfg.eval
_validate_parallel_eval_cfg(eval_cfg)
if _is_cuda_device(eval_cfg.get('device', 'cpu')):
_resolve_cuda_devices(eval_cfg)
return _run_eval_parallel_cuda(cfg)
return _run_eval_parallel_cpu(cfg)
def _run_eval_parallel_cpu(cfg: DictConfig):
eval_cfg = cfg.eval
artifact_paths = _resolve_artifact_paths(eval_cfg)
payloads, active_workers = _build_parallel_worker_payloads(cfg, artifact_paths)
try:
worker_summaries = _run_spawn_jobs(
payloads=payloads,
max_workers=max(active_workers, 1),
worker_fn=_run_eval_worker_entry,
)
except Exception as exc:
raise RuntimeError(f'Parallel rollout worker failed: {exc}') from exc
summary = _merge_worker_summaries(worker_summaries, artifact_paths)
public_summary = _public_summary(summary)
if artifact_paths.get('summary_json') is not None:
_save_summary_json(artifact_paths['summary_json'], public_summary)
if artifact_paths.get('timing_json') is not None:
_save_summary_json(artifact_paths['timing_json'], public_summary.get('timing_summary', {}))
return public_summary
def _run_eval_parallel_cuda(cfg: DictConfig):
eval_cfg = cfg.eval
_validate_parallel_eval_cfg(eval_cfg)
artifact_paths = _resolve_artifact_paths(eval_cfg)
worker_payloads, _active_workers = _build_parallel_worker_payloads(cfg, artifact_paths)
cuda_devices = _resolve_cuda_devices(eval_cfg)
server_payloads, assigned_worker_payloads = _build_cuda_server_payloads(
cfg,
worker_payloads=worker_payloads,
cuda_devices=cuda_devices,
)
try:
worker_summaries = _run_cuda_parallel_processes(
server_payloads=server_payloads,
worker_payloads=assigned_worker_payloads,
)
except Exception as exc:
raise RuntimeError(f'Parallel CUDA rollout failed: {exc}') from exc
summary = _merge_worker_summaries(worker_summaries, artifact_paths)
public_summary = _public_summary(summary)
if artifact_paths.get('summary_json') is not None:
_save_summary_json(artifact_paths['summary_json'], public_summary)
if artifact_paths.get('timing_json') is not None:
_save_summary_json(artifact_paths['timing_json'], public_summary.get('timing_summary', {}))
return public_summary
def _run_eval(cfg: DictConfig):
"""
使用 agent 内置队列管理的简化版 VLA 评估
所有评估参数来自 vla/conf/eval.yaml,合并到 cfg 中。
命令行覆盖: python eval_vla_simple.py eval.ckpt_path=... eval.num_episodes=5
"""
_print_eval_config(cfg)
requested_workers = int(cfg.eval.get('num_workers', 1))
active_workers = _normalize_num_workers(
num_workers=requested_workers,
num_episodes=int(cfg.eval.get('num_episodes', 0)),
)
if active_workers <= 1:
return _run_eval_serial(cfg)
return _run_eval_parallel(cfg)
@hydra.main(version_base=None, config_path="../../vla/conf", config_name="config")
def main(cfg: DictConfig):
return _run_eval(cfg)
if __name__ == '__main__':
main()