Files
VLA/wall_x/serving/client.py
T
Lufang Chen 722335e428 Enable inference serving && fix train stability (#59)
* enable serving

* lint

* update

* update
2025-10-26 16:02:34 +08:00

381 lines
12 KiB
Python

#!/usr/bin/env python3
"""
Example client for Wall-X model server with sync support.
This script demonstrates how to connect to a Wall-X server and request
action predictions from observations in both sync and async contexts.
"""
import asyncio
import logging
from typing import Dict, List
import numpy as np
import threading
import yaml
import torch
import matplotlib.pyplot as plt
import os
from wall_x.model.action_head import Normalizer
from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl_act import (
Qwen2_5_VLMoEForAction,
)
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
from wall_x.utils.constant import action_statistic_dof
try:
import msgpack
import msgpack_numpy as m
m.patch()
except ImportError:
print("Please install msgpack-numpy: pip install msgpack-numpy")
exit(1)
try:
import websockets
except ImportError:
print("Please install websockets: pip install websockets")
exit(1)
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class WallXClient:
"""Client for connecting to Wall-X model server."""
def __init__(self, config_path: str, uri: str = "ws://localhost:8000"):
"""Initialize client.
Args:
uri: WebSocket URI of the server (e.g., ws://localhost:8000)
"""
self.uri = uri
self.websocket = None
self.metadata = None
self._loop = None
self._thread = None
with open(config_path, "r") as f:
self.train_config = yaml.load(f, Loader=yaml.FullLoader)
self.init_normalizer(self.train_config)
async def connect(self):
"""Connect to the server and receive metadata."""
logger.info(f"Connecting to {self.uri}...")
self.websocket = await websockets.connect(
self.uri,
ping_interval=None,
ping_timeout=None,
max_size=None,
)
self.metadata = msgpack.unpackb(await self.websocket.recv())
logger.info(f"Connected! Server metadata: {self.metadata}")
async def predict(self, obs: Dict) -> Dict:
"""Get action prediction from observation.
Args:
obs: Observation dictionary containing:
- 'image': Image array (H, W, C)
- 'prompt': Optional text prompt
- 'state': Optional robot state
Returns:
Dictionary with:
- 'action': Predicted action array
- 'server_timing': Timing information
"""
if self.websocket is None:
raise RuntimeError("Not connected. Call connect() first.")
await self.websocket.send(msgpack.packb(obs))
response = msgpack.unpackb(await self.websocket.recv())
return response
async def close(self):
"""Close the connection."""
if self.websocket:
await self.websocket.close()
logger.info("Connection closed")
async def reset(self):
"""Reset the policy (if supported)."""
pass
# ============ Synchronous methods (using independent thread event loop) ============
def _start_background_loop(self):
"""Start event loop in background thread."""
self._loop = asyncio.new_event_loop()
asyncio.set_event_loop(self._loop)
self._loop.run_forever()
def _ensure_loop(self):
"""Ensure background event loop is running."""
if self._loop is None or not self._loop.is_running():
self._thread = threading.Thread(
target=self._start_background_loop, daemon=True
)
self._thread.start()
# Wait for loop to start
import time
while self._loop is None:
time.sleep(0.01)
def _run_async(self, coro):
"""Run coroutine in background event loop."""
self._ensure_loop()
future = asyncio.run_coroutine_threadsafe(coro, self._loop)
return future.result()
def connect_sync(self):
"""Synchronously connect to server."""
return self._run_async(self.connect())
def norm_state(
self,
state: np.ndarray,
dataset_names: List[str],
state_mask: torch.Tensor = None,
) -> np.ndarray:
"""Normalize state."""
return self.normalizer_propri.normalize_data(state, dataset_names, state_mask)
def predict_sync(self, obs: Dict) -> Dict:
"""Synchronous prediction method.
Args:
obs: Observation dictionary
Returns:
Prediction result dictionary
"""
return self._run_async(self.predict(obs))
def close_sync(self):
"""Synchronously close connection."""
result = self._run_async(self.close())
# Stop event loop
if self._loop:
self._loop.call_soon_threadsafe(self._loop.stop)
return result
def init_normalizer(self, train_config):
customized_dof_config = train_config["customized_robot_config"][
"customized_dof_config"
]
customized_agent_pos_config = train_config["customized_robot_config"][
"customized_agent_pos_config"
]
Qwen2_5_VLMoEForAction._set_customized_config(train_config)
self.normalizer_action = Normalizer(
action_statistic_dof, customized_dof_config
).to("cuda")
self.normalizer_propri = Normalizer(
action_statistic_dof, customized_agent_pos_config
).to("cuda")
print("Normalizer initialized")
def prepare_batch_sync(data, normalizer_action, normalizer_propri, dataset_names):
"""Synchronous version of prepare_batch."""
image = (data["image"].permute(1, 2, 0) * 255).to(torch.uint8).cpu().numpy()
wrist_image = (
(data["wrist_image"].permute(1, 2, 0) * 255).to(torch.uint8).cpu().numpy()
)
prompt = data["task"]
state = data["state"].to("cuda")
if state.dim() == 1:
state = state.unsqueeze(0)
state_mask = torch.ones([1, 32, 20]).to("cuda")
state_mask[:, :, 8:] = 0
state = normalizer_propri.normalize_data(state, dataset_names, state_mask)
state = state.cpu().numpy().astype(np.float32)
obs = {
"front_view": image,
"left_wrist_view": wrist_image,
"prompt": prompt,
"state": state,
"dataset_names": dataset_names,
}
return obs
def init_serving_sample_dataset(train_config):
repo_id = train_config["data"]["lerobot_config"]["repo_id"]
meta_info = LeRobotDatasetMetadata(repo_id)
dataset_fps = meta_info.fps
delta_timestamps = {
"actions": [t / dataset_fps for t in range(32)],
}
dataset = LeRobotDataset(
repo_id,
episodes=[0],
delta_timestamps=delta_timestamps,
video_backend="pyav",
)
return dataset, repo_id
# ============ Synchronous version of main function ============
def main_sync(args):
"""Synchronous version of main function."""
# Create client and connect
client = WallXClient(args.config_path, uri=args.uri)
client.connect_sync()
dataset, repo_id = init_serving_sample_dataset(client.train_config)
total_frames = len(dataset)
gt_traj = np.zeros((total_frames, args.action_dim))
pred_traj = np.zeros((total_frames, args.action_dim))
import torch
dof_mask = torch.ones([1, 32, 20]).to("cuda")
dof_mask[:, :, args.action_dim :] = 0
# Synchronous processing
for idx, data in enumerate(dataset):
if idx % args.pred_horizon == 0 and idx + args.pred_horizon < total_frames:
print(f"Processing frame {idx}")
obs = prepare_batch_sync(
data,
client.normalizer_action,
client.normalizer_propri,
dataset_names=[repo_id],
)
response = client.predict_sync(obs)
pred_action = response["action"]
pred_traj[idx : idx + args.pred_horizon] = pred_action
gt_traj[idx : idx + args.pred_horizon] = data["actions"]
# Draw plot
timesteps = gt_traj.shape[0]
fig, axs = plt.subplots(
args.action_dim, 1, figsize=(15, 5 * args.action_dim), sharex=True
)
fig.suptitle("Action Comparison for lerobot", fontsize=16)
for i in range(args.action_dim):
axs[i].plot(range(timesteps), gt_traj[:, i], label="Ground Truth")
axs[i].plot(range(timesteps), pred_traj[:, i], label="Prediction")
axs[i].set_ylabel(f"Action Dim {i+1}")
axs[i].legend()
axs[i].grid(True)
axs[-1].set_xlabel("Timestep")
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
os.makedirs(args.save_dir, exist_ok=True)
save_path = os.path.join(args.save_dir, "lerobot_comparison_serving.png")
plt.savefig(save_path)
print(f"Saved plot to {save_path}")
plt.close()
# Close connection
client.close_sync()
# ============ Asynchronous version of main function (keep original functionality) ============
async def main(args):
client = WallXClient(args.config_path, uri=args.uri)
await client.connect()
dataset, repo_id = init_serving_sample_dataset(client.train_config)
total_frames = len(dataset)
gt_traj = np.zeros((total_frames, args.action_dim))
pred_traj = np.zeros((total_frames, args.action_dim))
for idx, data in enumerate(dataset):
if idx % args.pred_horizon == 0 and idx + args.pred_horizon < total_frames:
print(f"Processing frame {idx}")
obs = prepare_batch_sync(
data,
client.normalizer_action,
client.normalizer_propri,
dataset_names=[repo_id],
)
response = await client.predict(obs)
pred_action = response["action"]
print(pred_action.shape)
pred_traj[idx : idx + args.pred_horizon] = pred_action
gt_traj[idx : idx + args.pred_horizon] = data["actions"]
timesteps = gt_traj.shape[0]
fig, axs = plt.subplots(
args.action_dim, 1, figsize=(15, 5 * args.action_dim), sharex=True
)
fig.suptitle("Action Comparison for lerobot", fontsize=16)
for i in range(args.action_dim):
axs[i].plot(range(timesteps), gt_traj[:, i], label="Ground Truth")
axs[i].plot(range(timesteps), pred_traj[:, i], label="Prediction")
axs[i].set_ylabel(f"Action Dim {i+1}")
axs[i].legend()
axs[i].grid(True)
axs[-1].set_xlabel("Timestep")
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
os.makedirs(args.save_dir, exist_ok=True)
save_path = os.path.join(args.save_dir, "lerobot_comparison_serving.png")
plt.savefig(save_path)
print(f"Saved plot to {save_path}")
plt.close()
if __name__ == "__main__":
"""Asynchronous version of main function."""
import argparse
parser = argparse.ArgumentParser(description="Wall-X client examples")
parser.add_argument(
"--example",
choices=["single", "multiple", "benchmark"],
default="single",
help="Example to run",
)
parser.add_argument(
"--uri",
default="ws://localhost:8000",
help="Server URI",
)
parser.add_argument(
"--pred_horizon", type=int, default=32, help="Prediction horizon"
)
parser.add_argument("--action_dim", type=int, default=7, help="Action dimension")
parser.add_argument(
"--config_path",
default="/x2robot_v2/vincent/workspace/opensource/cfg/config_from_qwen_libero.yml",
help="Train config path",
)
parser.add_argument(
"--save_dir",
default="/x2robot_v2/vincent/workspace/opensource/plots/libero",
help="Save directory",
)
args = parser.parse_args()
# Synchronous mode
main_sync(args)
# Asynchronous mode
# asyncio.run(main(args))