Enable inference serving && fix train stability (#59)
* enable serving * lint * update * update
This commit is contained in:
@@ -0,0 +1,380 @@
|
||||
#!/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))
|
||||
Reference in New Issue
Block a user