Enable inference serving && fix train stability (#59)
* enable serving * lint * update * update
This commit is contained in:
@@ -307,7 +307,7 @@ class ActionProcessor(nn.Module):
|
||||
torch.Tensor: Sampled timesteps of shape [batch_size]
|
||||
"""
|
||||
sample = self.beta_dist.sample([batch_size]).to(device=device, dtype=dtype)
|
||||
time = (self.s - sample) / self.s
|
||||
time = (1 - sample) / self.s
|
||||
return time
|
||||
|
||||
def proprioception_proj(
|
||||
|
||||
@@ -0,0 +1,263 @@
|
||||
# Wall-X Model Serving
|
||||
|
||||
This directory contains scripts for serving Wall-X models via a websocket server, allowing remote clients to connect and get action predictions from observations.
|
||||
|
||||
## Overview
|
||||
|
||||
The serving infrastructure consists of three main components:
|
||||
|
||||
1. **WebsocketPolicyServer** (`wall_x/serving/websocket_policy_server.py`): Generic websocket server that can serve any policy implementing the `BasePolicy` interface
|
||||
2. **WallXPolicy** (`wall_x/serving/policy/wall_x_policy.py`): Policy wrapper that adapts the Wall-X model to the `BasePolicy` interface
|
||||
3. **launch_serving.py**: Main script for starting the server with various configurations
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Basic Usage
|
||||
|
||||
Serve a model with default LIBERO configuration:
|
||||
|
||||
```bash
|
||||
cd /x2robot_v2/vincent/workspace/opensource
|
||||
python -m wall_x.serving.launch_serving \
|
||||
--env libero \
|
||||
--model-config.model-path /path/to/libero_model_stuff \
|
||||
--model-config.action-tokenizer-path /path/to/fast/ \
|
||||
--model-config.train-config-path /path/to/config.yml
|
||||
```
|
||||
|
||||
### Specify Environment
|
||||
|
||||
Serve with a specific environment preset:
|
||||
|
||||
```bash
|
||||
# LIBERO (single arm, 7 DOF)
|
||||
python -m wall_x.serving.launch_serving --env libero
|
||||
|
||||
# ALOHA (dual arm, 14 DOF)
|
||||
python -m wall_x.serving.launch_serving --env aloha
|
||||
```
|
||||
|
||||
### Custom Configuration
|
||||
|
||||
Serve with custom model paths and settings:
|
||||
|
||||
```bash
|
||||
python -m wall_x.serving.launch_serving \
|
||||
--model-config.model-path /path/to/model \
|
||||
--model-config.action-tokenizer-path /path/to/tokenizer \
|
||||
--model-config.train-config-path /path/to/train_config.yml \
|
||||
--model-config.action-dim 7 \
|
||||
--model-config.state-dim 8 \
|
||||
--model-config.pred-horizon 32 \
|
||||
--model-config.camera-key front_view left_wrist_view \
|
||||
--port 8000
|
||||
```
|
||||
|
||||
## Command Line Arguments
|
||||
|
||||
### Basic Arguments
|
||||
|
||||
- `--env {libero,aloha}`: Environment mode (default: libero)
|
||||
- `--port PORT`: Port to serve on (default: 8000)
|
||||
- `--host HOST`: Host to bind to (default: 0.0.0.0)
|
||||
- `--default-prompt TEXT`: Default text prompt if not provided in observation
|
||||
- `--debug`: Enable debug logging
|
||||
|
||||
### Model Configuration
|
||||
|
||||
All model configuration arguments use the `--model-config.` prefix:
|
||||
|
||||
- `--model-config.model-path PATH`: Path to pretrained model checkpoint (required)
|
||||
- `--model-config.action-tokenizer-path PATH`: Path to action tokenizer (required)
|
||||
- `--model-config.train-config-path PATH`: Path to train config YAML file (required)
|
||||
- `--model-config.action-dim INT`: Action space dimension (default: 7)
|
||||
- `--model-config.state-dim INT`: Robot state dimension (default: 8)
|
||||
- `--model-config.pred-horizon INT`: Prediction horizon (default: 32)
|
||||
- `--model-config.device {cuda,cpu}`: Device to run on (default: cuda)
|
||||
- `--model-config.dtype {bfloat16,float16,float32}`: Model dtype (default: bfloat16)
|
||||
- `--model-config.predict-mode {fast,diffusion}`: Prediction mode (default: fast)
|
||||
- `--model-config.camera-key KEY1 KEY2 ...`: Camera keys for observation images
|
||||
|
||||
### Camera Keys
|
||||
|
||||
The `camera-key` parameter specifies which camera views are expected in the observation dictionary. This is **critical** for proper operation:
|
||||
|
||||
- Keys must match between server configuration and client observations
|
||||
- Order matters: keys are processed in the order specified
|
||||
- Common keys: `front_view`, `left_wrist_view`, `right_wrist_view`, `face_view`
|
||||
|
||||
Example:
|
||||
```bash
|
||||
--model-config.camera-key front_view left_wrist_view
|
||||
```
|
||||
|
||||
Client must send observations with matching keys:
|
||||
```python
|
||||
obs = {
|
||||
"front_view": image1, # Must match camera-key[0]
|
||||
"left_wrist_view": image2, # Must match camera-key[1]
|
||||
"prompt": "task description",
|
||||
"state": robot_state,
|
||||
}
|
||||
```
|
||||
|
||||
## Default Configurations
|
||||
|
||||
### LIBERO (Single Arm)
|
||||
|
||||
```python
|
||||
ModelConfig(
|
||||
model_path="/path/to/model",
|
||||
action_tokenizer_path="/path/to/action_tokenizer",
|
||||
train_config_path="/path/to/train_config",
|
||||
state_dim=8,
|
||||
action_dim=7,
|
||||
pred_horizon=32,
|
||||
device="cuda",
|
||||
dtype="bfloat16",
|
||||
predict_mode="fast",
|
||||
camera_key=["front_view", "left_wrist_view"],
|
||||
)
|
||||
```
|
||||
|
||||
### ALOHA (Dual Arm)
|
||||
|
||||
```python
|
||||
ModelConfig(
|
||||
model_path="/path/to/model",
|
||||
action_tokenizer_path="/path/to/action_tokenizer",
|
||||
train_config_path="/path/to/train_config",
|
||||
state_dim=14,
|
||||
action_dim=14,
|
||||
pred_horizon=32,
|
||||
device="cuda",
|
||||
dtype="bfloat16",
|
||||
predict_mode="fast",
|
||||
camera_key=["face_view", "left_wrist_view", "right_wrist_view"],
|
||||
)
|
||||
```
|
||||
|
||||
## Server Protocol
|
||||
|
||||
### Connection Flow
|
||||
|
||||
1. Client connects to `ws://host:port`
|
||||
2. Server sends metadata JSON with policy information
|
||||
3. Client sends observation (msgpack-encoded)
|
||||
4. Server responds with action prediction (msgpack-encoded)
|
||||
5. Repeat steps 3-4 for each inference
|
||||
|
||||
### Observation Format
|
||||
|
||||
Observations must be a dictionary with camera keys matching server configuration:
|
||||
|
||||
```python
|
||||
obs = {
|
||||
# Image observations - keys must match server's camera_key configuration
|
||||
"front_view": np.ndarray, # (H, W, 3) uint8 or float
|
||||
"left_wrist_view": np.ndarray, # (H, W, 3) uint8 or float
|
||||
|
||||
# Required fields
|
||||
"prompt": str, # Task description
|
||||
"dataset_names": List[str], # Dataset/robot name, e.g., ["physical-intelligence/libero"]
|
||||
"state": np.ndarray, # Robot proprioception state (state_dim,)
|
||||
}
|
||||
```
|
||||
|
||||
**Important**: The image keys (`front_view`, `left_wrist_view`, etc.) must exactly match the `camera_key` parameter configured on the server.
|
||||
|
||||
### Action Response Format
|
||||
|
||||
Actions are returned as a dictionary:
|
||||
|
||||
```python
|
||||
{
|
||||
"action": np.ndarray, # Predicted action [pred_horizon, action_dim]
|
||||
"server_timing": {
|
||||
"infer_ms": float, # Inference time in milliseconds
|
||||
"prev_total_ms": float, # Total time for previous request
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Server Metadata
|
||||
|
||||
When connecting, the server sends metadata:
|
||||
|
||||
```python
|
||||
{
|
||||
"action_dim": int, # Action space dimension
|
||||
"pred_horizon": int, # Number of future actions predicted
|
||||
"device": str, # Device model runs on
|
||||
"predict_mode": str, # Prediction mode (fast/diffusion)
|
||||
"env": str, # Environment name
|
||||
}
|
||||
```
|
||||
|
||||
### Health Check
|
||||
|
||||
HTTP health check endpoint available at:
|
||||
```
|
||||
http://host:port/healthz
|
||||
```
|
||||
|
||||
Returns `200 OK` if the server is running.
|
||||
|
||||
## Client Example
|
||||
|
||||
### Synchronous Python Client
|
||||
|
||||
For synchronous usage, see `wall_x/serving/client.py`:
|
||||
|
||||
```python
|
||||
from wall_x.serving.client import WallXClient
|
||||
|
||||
# Create and connect
|
||||
client = WallXClient(uri="ws://localhost:8000")
|
||||
client.connect_sync()
|
||||
|
||||
# Prepare observation
|
||||
obs = {
|
||||
"front_view": image1,
|
||||
"left_wrist_view": image2,
|
||||
"prompt": "task description",
|
||||
"state": robot_state,
|
||||
"dataset_names": ["physical-intelligence/libero"],
|
||||
}
|
||||
|
||||
# Get prediction
|
||||
response = client.predict_sync(obs)
|
||||
action = response["action"]
|
||||
|
||||
# Close connection
|
||||
client.close_sync()
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
### WebsocketPolicyServer
|
||||
|
||||
Generic websocket server that:
|
||||
- Handles websocket connections with msgpack serialization
|
||||
- Tracks inference timing and performance metrics
|
||||
- Provides health check endpoint
|
||||
- Handles errors gracefully with proper logging
|
||||
- Supports concurrent client connections
|
||||
|
||||
### WallXPolicy
|
||||
|
||||
Policy wrapper that:
|
||||
- Loads and manages the Wall-X model from pretrained checkpoint
|
||||
- Processes multi-camera observations
|
||||
- Handles image preprocessing (smart resize, normalization)
|
||||
- Manages device placement and dtype conversion
|
||||
- Provides policy metadata to clients
|
||||
- Supports both fast tokenizer and diffusion prediction modes
|
||||
|
||||
### Image Processing Pipeline
|
||||
|
||||
1. **Camera Key Matching**: Extracts images from observation dict using configured camera keys
|
||||
2. **Format Conversion**: Converts numpy arrays to PIL Images
|
||||
3. **Smart Resize**: Applies Qwen's smart resize algorithm based on min/max pixels
|
||||
4. **Vision Token Formatting**: Inserts vision tokens in text prompt
|
||||
5. **Batch Preparation**: Creates model-ready BatchFeature input
|
||||
@@ -0,0 +1,3 @@
|
||||
from .websocket_policy_server import WebsocketPolicyServer, BasePolicy
|
||||
|
||||
__all__ = ["WebsocketPolicyServer", "BasePolicy"]
|
||||
@@ -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))
|
||||
@@ -0,0 +1,220 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Server script for Wall-X model.
|
||||
|
||||
This script serves a Wall-X model using a websocket server, allowing
|
||||
clients to connect and get action predictions from observations.
|
||||
|
||||
Based on the OpenPI serve_policy.py script structure.
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import field
|
||||
import enum
|
||||
import logging
|
||||
import socket
|
||||
import sys
|
||||
import yaml
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import tyro
|
||||
|
||||
from wall_x.serving.policy.wall_x_policy import WallXPolicy
|
||||
from wall_x.serving.websocket_policy_server import WebsocketPolicyServer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class EnvMode(enum.Enum):
|
||||
"""Supported environments/datasets."""
|
||||
|
||||
LIBERO = "libero"
|
||||
ALOHA = "aloha"
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ModelConfig:
|
||||
"""Configuration for loading a Wall-X model."""
|
||||
|
||||
# Path to the pretrained model checkpoint
|
||||
model_path: str
|
||||
# Path to the action tokenizer
|
||||
action_tokenizer_path: str
|
||||
# Path to train config yaml
|
||||
train_config_path: str
|
||||
# Action dimension for the environment
|
||||
action_dim: int = 7
|
||||
# State dimension for the environment
|
||||
state_dim: int = 8
|
||||
# Prediction horizon (number of future actions to predict)
|
||||
pred_horizon: int = 32
|
||||
# Device to run model on
|
||||
device: str = "cuda"
|
||||
# Model dtype (bfloat16, float16, float32)
|
||||
dtype: str = "bfloat16"
|
||||
# Prediction mode (fast or slow)
|
||||
predict_mode: str = "fast"
|
||||
# Camera key for the environment
|
||||
camera_key: List[str] = field(
|
||||
default_factory=lambda: ["front_view", "left_wrist_view", "right_wrist_view"]
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class Args:
|
||||
"""Arguments for the serve_wall_x script."""
|
||||
|
||||
# Environment mode (used for default configurations)
|
||||
env: EnvMode = EnvMode.LIBERO
|
||||
|
||||
# Model configuration. If not provided, uses default config for the environment
|
||||
model_config: ModelConfig | None = None
|
||||
|
||||
# Default text prompt to use if not provided in observation
|
||||
default_prompt: str | None = None
|
||||
|
||||
# Port to serve the policy on
|
||||
port: int = 8000
|
||||
|
||||
# Host to bind the server to
|
||||
host: str = "0.0.0.0"
|
||||
|
||||
# Enable debug logging
|
||||
debug: bool = False
|
||||
|
||||
|
||||
# Default model configurations for each environment
|
||||
DEFAULT_CONFIGS: dict[EnvMode, ModelConfig] = {
|
||||
EnvMode.LIBERO: ModelConfig(
|
||||
model_path="/path/to/model",
|
||||
action_tokenizer_path="/path/to/action_tokenizer",
|
||||
train_config_path="/path/to/train_config",
|
||||
state_dim=8,
|
||||
action_dim=7,
|
||||
pred_horizon=32,
|
||||
device="cuda",
|
||||
dtype="bfloat16",
|
||||
predict_mode="fast",
|
||||
camera_key=["front_view", "left_wrist_view"],
|
||||
),
|
||||
EnvMode.ALOHA: ModelConfig(
|
||||
model_path="/path/to/model",
|
||||
action_tokenizer_path="/path/to/action_tokenizer",
|
||||
train_config_path="/path/to/train_config",
|
||||
state_dim=14,
|
||||
action_dim=14,
|
||||
pred_horizon=32,
|
||||
device="cuda",
|
||||
dtype="bfloat16",
|
||||
predict_mode="fast",
|
||||
camera_key=["face_view", "left_wrist_view", "right_wrist_view"],
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def get_model_config(args: Args) -> ModelConfig:
|
||||
"""Get model configuration from args or defaults."""
|
||||
if args.model_config is not None:
|
||||
return args.model_config
|
||||
|
||||
if config := DEFAULT_CONFIGS.get(args.env):
|
||||
logger.info(f"Using default configuration for {args.env.value}")
|
||||
return config
|
||||
|
||||
raise ValueError(
|
||||
f"No default configuration for {args.env.value}. "
|
||||
f"Please provide --model-config with model_path and action_tokenizer_path."
|
||||
)
|
||||
|
||||
|
||||
def create_policy(args: Args) -> WallXPolicy:
|
||||
"""Create a Wall-X policy from the given arguments."""
|
||||
config = get_model_config(args)
|
||||
logger.info(f"Creating Wall-X policy with config: {config}")
|
||||
|
||||
# Validate paths
|
||||
if not Path(config.model_path).exists():
|
||||
logger.warning(f"Model path does not exist: {config.model_path}")
|
||||
|
||||
if not Path(config.action_tokenizer_path).exists():
|
||||
logger.warning(
|
||||
f"Action tokenizer path does not exist: {config.action_tokenizer_path}"
|
||||
)
|
||||
|
||||
with open(config.train_config_path, "r") as f:
|
||||
train_config = yaml.load(f, Loader=yaml.FullLoader)
|
||||
|
||||
policy = WallXPolicy(
|
||||
model_path=config.model_path,
|
||||
train_config=train_config,
|
||||
action_tokenizer_path=config.action_tokenizer_path,
|
||||
action_dim=config.action_dim,
|
||||
agent_pos_dim=config.state_dim,
|
||||
pred_horizon=config.pred_horizon,
|
||||
device=config.device,
|
||||
dtype=config.dtype,
|
||||
predict_mode=config.predict_mode,
|
||||
default_prompt=args.default_prompt,
|
||||
camera_key=config.camera_key,
|
||||
)
|
||||
|
||||
return policy
|
||||
|
||||
|
||||
def main(args: Args) -> None:
|
||||
"""Main function to start the Wall-X model server."""
|
||||
log_level = logging.DEBUG if args.debug else logging.INFO
|
||||
logging.basicConfig(
|
||||
level=log_level,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
)
|
||||
|
||||
logger.info("Starting Wall-X model server")
|
||||
logger.info(f"Environment: {args.env.value}")
|
||||
logger.info(f"Port: {args.port}")
|
||||
logger.info(f"Host: {args.host}")
|
||||
|
||||
# Create policy
|
||||
try:
|
||||
policy = create_policy(args)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to create policy: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# Get policy metadata
|
||||
policy_metadata = policy.metadata
|
||||
policy_metadata["env"] = args.env.value
|
||||
|
||||
# Get network info
|
||||
hostname = socket.gethostname()
|
||||
try:
|
||||
local_ip = socket.gethostbyname(hostname)
|
||||
except Exception:
|
||||
local_ip = "unknown"
|
||||
|
||||
logger.info(f"Server hostname: {hostname}")
|
||||
logger.info(f"Server IP: {local_ip}")
|
||||
logger.info(f"Server will be available at: ws://{args.host}:{args.port}")
|
||||
logger.info(f"Health check endpoint: http://{args.host}:{args.port}/healthz")
|
||||
|
||||
# Create and start server
|
||||
server = WebsocketPolicyServer(
|
||||
policy=policy,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
metadata=policy_metadata,
|
||||
)
|
||||
|
||||
logger.info("Starting server...")
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
logger.info("Server stopped by user")
|
||||
except Exception as e:
|
||||
logger.error(f"Server error: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main(tyro.cli(Args))
|
||||
@@ -0,0 +1,3 @@
|
||||
from .wall_x_policy import WallXPolicy
|
||||
|
||||
__all__ = ["WallXPolicy"]
|
||||
@@ -0,0 +1,246 @@
|
||||
from typing import Dict, List
|
||||
import logging
|
||||
import numpy as np
|
||||
from wall_x.data.utils import preprocesser_call
|
||||
from qwen_vl_utils.vision_process import smart_resize
|
||||
import torch
|
||||
from PIL import Image
|
||||
from transformers import BatchFeature
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def prepare_batch(
|
||||
obs: Dict,
|
||||
processor,
|
||||
camera_key: List[str],
|
||||
agent_pos_dim,
|
||||
action_dim,
|
||||
pred_horizon,
|
||||
fixed_action_dim,
|
||||
max_length,
|
||||
image_factor: int,
|
||||
min_pixels: int,
|
||||
max_pixels: int,
|
||||
predict_mode: str = "fast",
|
||||
device: str = "cuda",
|
||||
) -> BatchFeature:
|
||||
"""Prepare observation into model input format.
|
||||
|
||||
Args:
|
||||
obs: Dictionary containing:
|
||||
- 'camera_key_0' : image 0
|
||||
- 'camera_key_1' : image 1
|
||||
...
|
||||
- 'prompt': Text prompt
|
||||
- 'state': Robot state/proprioception
|
||||
- 'dataset_names': Dataset names
|
||||
|
||||
Returns:
|
||||
BatchFeature object ready for model input
|
||||
"""
|
||||
# Handle images - can be single image, list of images, or dict of images
|
||||
images = []
|
||||
images = [obs[key] for key in camera_key]
|
||||
# Convert numpy arrays to PIL Images
|
||||
processed_images = []
|
||||
for img in images:
|
||||
if isinstance(img, np.ndarray):
|
||||
# Debug: Log the shape and dtype
|
||||
logger.debug(f"Image shape: {img.shape}, dtype: {img.dtype}")
|
||||
|
||||
# Handle unexpected dimensions - squeeze if needed
|
||||
if img.ndim > 3:
|
||||
logger.warning(
|
||||
f"Image has {img.ndim} dimensions, squeezing extra dimensions"
|
||||
)
|
||||
img = np.squeeze(img)
|
||||
|
||||
# Verify shape is valid for PIL
|
||||
if img.ndim == 2:
|
||||
# Grayscale image
|
||||
pass
|
||||
elif img.ndim == 3:
|
||||
# Check if channel dimension is first or last
|
||||
if img.shape[0] == 3 or img.shape[0] == 1:
|
||||
# Channels first, transpose to channels last
|
||||
img = np.transpose(img, (1, 2, 0))
|
||||
elif img.shape[2] == 3 or img.shape[2] == 1:
|
||||
# Already channels last
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected image shape: {img.shape}. Expected (H, W, C) or (C, H, W)"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid image dimensions: {img.ndim}. Expected 2 or 3 dimensions, got shape {img.shape}"
|
||||
)
|
||||
|
||||
# Convert to PIL Image
|
||||
if img.dtype == np.uint8:
|
||||
img = Image.fromarray(img)
|
||||
else:
|
||||
img = Image.fromarray((img * 255).astype(np.uint8))
|
||||
processed_images.append(img)
|
||||
|
||||
# Apply smart resize to images
|
||||
resized_images = process_images(
|
||||
processed_images, image_factor, min_pixels, max_pixels
|
||||
)
|
||||
|
||||
# Handle text prompt - format with vision tokens
|
||||
instruction = obs["prompt"]
|
||||
formatted_text = format_text_with_vision_tokens(
|
||||
instruction, camera_key, predict_mode, pred_horizon
|
||||
)
|
||||
|
||||
# Use processor to prepare inputs
|
||||
inputs = preprocesser_call(
|
||||
processor=processor,
|
||||
text=[formatted_text],
|
||||
images=[resized_images],
|
||||
videos=None,
|
||||
padding=True,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
max_length=max_length,
|
||||
)
|
||||
|
||||
action_token_id = processor.tokenizer.convert_tokens_to_ids("<|action|>")
|
||||
moe_token_types = inputs.input_ids == action_token_id
|
||||
inputs["moe_token_types"] = moe_token_types
|
||||
|
||||
# Handle robot state/proprioception if available
|
||||
if "state" in obs:
|
||||
state = obs["state"]
|
||||
if isinstance(state, np.ndarray):
|
||||
state = torch.from_numpy(state).float()
|
||||
elif not isinstance(state, torch.Tensor):
|
||||
state = torch.tensor(state, dtype=torch.float32)
|
||||
|
||||
# Add batch dimension if needed
|
||||
if state.dim() == 1:
|
||||
state = state.unsqueeze(0)
|
||||
if state.dim() == 2:
|
||||
state = state.unsqueeze(1) # [batch, 1, state_dim]
|
||||
|
||||
# Pad to 20 dimensions if needed (same as training)
|
||||
if state.shape[-1] < 20:
|
||||
padding = torch.zeros(state.shape[0], state.shape[1], 20 - state.shape[-1])
|
||||
state = torch.cat([state, padding], dim=-1)
|
||||
|
||||
# Create mask for valid dimensions
|
||||
agent_pos_mask = torch.ones_like(state)
|
||||
if state.shape[-1] > agent_pos_dim:
|
||||
agent_pos_mask[:, :, agent_pos_dim:] = 0
|
||||
|
||||
inputs["proprioception"] = state
|
||||
inputs["agent_pos_mask"] = agent_pos_mask
|
||||
|
||||
# Add dataset name (required by model)
|
||||
inputs["dataset_names"] = obs["dataset_names"]
|
||||
|
||||
# Move all tensors to device
|
||||
for key in inputs:
|
||||
if isinstance(inputs[key], torch.Tensor):
|
||||
inputs[key] = inputs[key].to(device)
|
||||
|
||||
dof_mask = torch.ones([state.shape[0], pred_horizon, fixed_action_dim])
|
||||
dof_mask[:, :, action_dim:] = 0
|
||||
|
||||
inputs["dof_mask"] = dof_mask
|
||||
|
||||
# Convert to BatchFeature to maintain consistency with training pipeline
|
||||
return BatchFeature(data=dict(inputs)).to(device)
|
||||
|
||||
|
||||
def process_images(
|
||||
images: List[Image.Image], image_factor: int, min_pixels: int, max_pixels: int
|
||||
) -> List[Image.Image]:
|
||||
"""Process images with smart resize following the data loading pattern.
|
||||
|
||||
Args:
|
||||
images: List of PIL Images
|
||||
|
||||
Returns:
|
||||
List of resized PIL Images
|
||||
"""
|
||||
resized_images = []
|
||||
for img_pil in images:
|
||||
current_width, current_height = img_pil.size
|
||||
|
||||
# Apply smart scaling (Qwen logic)
|
||||
resized_height, resized_width = smart_resize(
|
||||
current_height,
|
||||
current_width,
|
||||
factor=image_factor,
|
||||
min_pixels=min_pixels,
|
||||
max_pixels=max_pixels,
|
||||
)
|
||||
|
||||
resized_img = img_pil.resize((resized_width, resized_height))
|
||||
resized_images.append(resized_img)
|
||||
|
||||
return resized_images
|
||||
|
||||
|
||||
def format_text_with_vision_tokens(
|
||||
instruction: str,
|
||||
camera_key: List[str],
|
||||
predict_mode: str = "fast",
|
||||
pred_horizon: int = 32,
|
||||
) -> str:
|
||||
"""Format text prompt with vision tokens for the model.
|
||||
|
||||
Args:
|
||||
instruction: Task instruction text
|
||||
camera_key: List of camera names
|
||||
|
||||
Returns:
|
||||
Formatted text with special tokens
|
||||
"""
|
||||
# Special tokens for formatting
|
||||
role_start_symbol = "<|im_start|>"
|
||||
role_end_symbol = "<|im_end|>"
|
||||
vision_start_symbol = "<|vision_start|>"
|
||||
vision_end_symbol = "<|vision_end|>"
|
||||
image_pad_symbol = "<|image_pad|>"
|
||||
propri_symbol = "<|propri|>"
|
||||
action_symbol = "<|action|>"
|
||||
# action_fast_symbol = "<|action_fast|>"
|
||||
|
||||
# Camera name mapping
|
||||
camera_name_mapping = {
|
||||
"front_view": "front view",
|
||||
"face_view": "front view",
|
||||
"left_wrist_view": "left wrist view",
|
||||
"right_wrist_view": "right wrist view",
|
||||
"top_view": "top view",
|
||||
"wall_view": "wall view",
|
||||
}
|
||||
pred_horizon = 32
|
||||
|
||||
# System prologue
|
||||
prologue = (
|
||||
f"{role_start_symbol}system\nYou are a helpful assistant.{role_end_symbol}\n"
|
||||
)
|
||||
|
||||
# User request with observation
|
||||
user_request = f"{role_start_symbol}user\nObservation:"
|
||||
if camera_key:
|
||||
for cam_name in camera_key:
|
||||
view_name = camera_name_mapping.get(cam_name, cam_name)
|
||||
user_request += f" {view_name}: {vision_start_symbol}{image_pad_symbol}{vision_end_symbol}"
|
||||
user_request += "\nInstruction:"
|
||||
|
||||
text_prompt = (
|
||||
f"\nPredict the next action in robot action.\nProprioception: {propri_symbol}\n"
|
||||
)
|
||||
user_message = f"{user_request} {instruction}{text_prompt}{role_end_symbol}\n"
|
||||
assistant_output = f"{role_start_symbol}assistant\n"
|
||||
if predict_mode == "diffusion":
|
||||
assistant_output += f"{action_symbol * pred_horizon}"
|
||||
complete_text = prologue + user_message + assistant_output
|
||||
|
||||
return complete_text
|
||||
@@ -0,0 +1,173 @@
|
||||
import logging
|
||||
from typing import Dict, Any, List
|
||||
import torch
|
||||
import numpy as np
|
||||
from transformers import AutoProcessor
|
||||
|
||||
from wall_x.serving.websocket_policy_server import BasePolicy
|
||||
from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl_act import Qwen2_5_VLMoEForAction
|
||||
from wall_x.serving.policy.utils import prepare_batch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class WallXPolicy(BasePolicy):
|
||||
"""Policy wrapper for Wall-X model that implements the BasePolicy interface."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_path: str,
|
||||
train_config: dict,
|
||||
action_tokenizer_path: str,
|
||||
action_dim: int,
|
||||
agent_pos_dim: int,
|
||||
pred_horizon: int,
|
||||
camera_key: List[str],
|
||||
device: str = "cuda",
|
||||
dtype: str = "bfloat16",
|
||||
predict_mode: str = "fast",
|
||||
default_prompt: str | None = None,
|
||||
min_pixels: int = 4 * 28 * 28,
|
||||
max_pixels: int = 16384 * 28 * 28,
|
||||
image_factor: int = 28,
|
||||
max_length: int = 768,
|
||||
):
|
||||
"""Initialize the Wall-X policy.
|
||||
|
||||
Args:
|
||||
model_path: Path to the pretrained model checkpoint
|
||||
action_tokenizer_path: Path to the action tokenizer
|
||||
action_dim: Dimension of action space
|
||||
pred_horizon: Prediction horizon for actions
|
||||
device: Device to run model on ('cuda' or 'cpu')
|
||||
dtype: Data type for model ('bfloat16', 'float16', or 'float32')
|
||||
predict_mode: Prediction mode ('fast' or 'slow')
|
||||
default_prompt: Default text prompt for the model
|
||||
min_pixels: Minimum pixels for image resizing
|
||||
max_pixels: Maximum pixels for image resizing
|
||||
image_factor: Factor for smart resize
|
||||
max_length: Maximum sequence length for text
|
||||
"""
|
||||
logger.info(f"Loading Wall-X model from {model_path}")
|
||||
|
||||
self.model = Qwen2_5_VLMoEForAction.from_pretrained(
|
||||
model_path,
|
||||
train_config=train_config,
|
||||
action_tokenizer_path=action_tokenizer_path,
|
||||
)
|
||||
self.model.eval()
|
||||
self.model = self.model.to(device)
|
||||
|
||||
self.model = self.model.bfloat16()
|
||||
|
||||
# hard code the action dim to 20 for align to wall-x configuration
|
||||
self.fixed_action_dim = 20
|
||||
|
||||
self.action_dim = action_dim
|
||||
self.agent_pos_dim = agent_pos_dim
|
||||
self.pred_horizon = pred_horizon
|
||||
self.device = device
|
||||
self.predict_mode = predict_mode
|
||||
self.default_prompt = default_prompt
|
||||
self.camera_key = camera_key
|
||||
|
||||
# Image preprocessing config
|
||||
self.min_pixels = min_pixels
|
||||
self.max_pixels = max_pixels
|
||||
self.image_factor = image_factor
|
||||
self.max_length = max_length
|
||||
|
||||
# Load processor
|
||||
logger.info("Loading processor and tokenizer...")
|
||||
self.processor = AutoProcessor.from_pretrained(model_path, use_fast=True)
|
||||
self.processor.tokenizer.padding_side = "left"
|
||||
|
||||
# Action buffer for multi-step predictions
|
||||
self.action_buffer = []
|
||||
self.buffer_index = 0
|
||||
|
||||
logger.info(
|
||||
f"Model loaded successfully. Device: {device}, Action dim: {action_dim}, Horizon: {pred_horizon}"
|
||||
)
|
||||
|
||||
@property
|
||||
def metadata(self) -> Dict[str, Any]:
|
||||
"""Return metadata about the policy."""
|
||||
return {
|
||||
"action_dim": self.action_dim,
|
||||
"pred_horizon": self.pred_horizon,
|
||||
"device": self.device,
|
||||
"predict_mode": self.predict_mode,
|
||||
}
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Reset the policy state."""
|
||||
self.action_buffer = []
|
||||
self.buffer_index = 0
|
||||
logger.debug("Policy reset")
|
||||
|
||||
def infer(self, obs: Dict) -> Dict:
|
||||
"""Infer action from observation.
|
||||
|
||||
Args:
|
||||
obs: Dictionary containing:
|
||||
- 'image': Image observation (numpy array or PIL Image)
|
||||
- 'prompt': Optional text prompt
|
||||
- 'state': Optional robot state
|
||||
- Other modality-specific observations
|
||||
|
||||
Returns:
|
||||
Dictionary containing:
|
||||
- 'action': Predicted action (numpy array)
|
||||
- Additional metadata
|
||||
"""
|
||||
try:
|
||||
# Need to predict new actions
|
||||
input_batch = prepare_batch(
|
||||
obs,
|
||||
self.processor,
|
||||
self.camera_key,
|
||||
self.agent_pos_dim,
|
||||
self.action_dim,
|
||||
self.pred_horizon,
|
||||
self.fixed_action_dim,
|
||||
self.max_length,
|
||||
self.image_factor,
|
||||
self.min_pixels,
|
||||
self.max_pixels,
|
||||
self.predict_mode,
|
||||
self.device,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = self.model(
|
||||
**input_batch,
|
||||
action_dim=(
|
||||
self.action_dim
|
||||
if self.predict_mode == "fast"
|
||||
else self.fixed_action_dim
|
||||
),
|
||||
pred_horizon=self.pred_horizon,
|
||||
mode="predict",
|
||||
predict_mode=self.predict_mode,
|
||||
)
|
||||
|
||||
if outputs["predict_action"] is None:
|
||||
predicted_actions = np.zeros(
|
||||
[1, self.pred_horizon, self.action_dim]
|
||||
).astype(np.float32)
|
||||
|
||||
predicted_actions = (
|
||||
outputs["predict_action"][:, :, : self.action_dim]
|
||||
.detach()
|
||||
.cpu()
|
||||
.to(torch.float32)
|
||||
.numpy()
|
||||
)
|
||||
|
||||
print(predicted_actions.shape)
|
||||
return {"action": predicted_actions}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error during inference: {e}")
|
||||
raise
|
||||
@@ -0,0 +1,132 @@
|
||||
import asyncio
|
||||
import http
|
||||
import logging
|
||||
import time
|
||||
import traceback
|
||||
from typing import Dict, Any
|
||||
|
||||
try:
|
||||
import msgpack
|
||||
import msgpack_numpy as m
|
||||
|
||||
m.patch()
|
||||
except ImportError:
|
||||
logging.warning(
|
||||
"msgpack-numpy not installed. Install with: pip install msgpack-numpy"
|
||||
)
|
||||
msgpack = None
|
||||
|
||||
import websockets.asyncio.server as _server
|
||||
import websockets.frames
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BasePolicy:
|
||||
"""Base class for policies that can be served."""
|
||||
|
||||
def infer(self, obs: Dict) -> Dict:
|
||||
"""Infer actions from observations."""
|
||||
raise NotImplementedError
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Reset the policy to its initial state."""
|
||||
pass
|
||||
|
||||
@property
|
||||
def metadata(self) -> Dict[str, Any]:
|
||||
"""Return metadata about the policy."""
|
||||
return {}
|
||||
|
||||
|
||||
class WebsocketPolicyServer:
|
||||
"""Serves a policy using the websocket protocol.
|
||||
|
||||
Implements a websocket server that:
|
||||
1. Sends policy metadata on connection
|
||||
2. Receives observations
|
||||
3. Returns predicted actions
|
||||
4. Tracks timing information
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
policy: BasePolicy,
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
metadata: Dict | None = None,
|
||||
) -> None:
|
||||
self._policy = policy
|
||||
self._host = host
|
||||
self._port = port
|
||||
self._metadata = metadata or {}
|
||||
logging.getLogger("websockets.server").setLevel(logging.INFO)
|
||||
|
||||
def serve_forever(self) -> None:
|
||||
asyncio.run(self.run())
|
||||
|
||||
async def run(self):
|
||||
async with _server.serve(
|
||||
self._handler,
|
||||
self._host,
|
||||
self._port,
|
||||
compression=None,
|
||||
max_size=None,
|
||||
ping_interval=None, # Disable automatic ping for long-running inference
|
||||
ping_timeout=None, # Disable ping timeout
|
||||
process_request=_health_check,
|
||||
) as server:
|
||||
logger.info(f"Server started on {self._host}:{self._port}")
|
||||
await server.serve_forever()
|
||||
|
||||
async def _handler(self, websocket: _server.ServerConnection):
|
||||
logger.info(f"Connection from {websocket.remote_address} opened")
|
||||
|
||||
if msgpack is None:
|
||||
await websocket.close(
|
||||
code=websockets.frames.CloseCode.INTERNAL_ERROR,
|
||||
reason="msgpack-numpy not installed on server",
|
||||
)
|
||||
return
|
||||
|
||||
# Send metadata to client
|
||||
await websocket.send(msgpack.packb(self._metadata))
|
||||
|
||||
prev_total_time = None
|
||||
while True:
|
||||
try:
|
||||
start_time = time.monotonic()
|
||||
obs = msgpack.unpackb(await websocket.recv())
|
||||
|
||||
infer_time = time.monotonic()
|
||||
action = self._policy.infer(obs)
|
||||
infer_time = time.monotonic() - infer_time
|
||||
|
||||
action["server_timing"] = {
|
||||
"infer_ms": infer_time * 1000,
|
||||
}
|
||||
if prev_total_time is not None:
|
||||
action["server_timing"]["prev_total_ms"] = prev_total_time * 1000
|
||||
|
||||
await websocket.send(msgpack.packb(action))
|
||||
prev_total_time = time.monotonic() - start_time
|
||||
|
||||
except websockets.ConnectionClosed:
|
||||
logger.info(f"Connection from {websocket.remote_address} closed")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error handling request: {e}")
|
||||
await websocket.send(traceback.format_exc())
|
||||
await websocket.close(
|
||||
code=websockets.frames.CloseCode.INTERNAL_ERROR,
|
||||
reason="Internal server error. Traceback included in previous frame.",
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def _health_check(
|
||||
connection: _server.ServerConnection, request: _server.Request
|
||||
) -> _server.Response | None:
|
||||
if request.path == "/healthz":
|
||||
return connection.respond(http.HTTPStatus.OK, "OK\n")
|
||||
return None
|
||||
Reference in New Issue
Block a user