[Serving] Update the details of the inference service (#70)

* Update the details of the inference service startup

* add functions to compute and update robot-specific action/state normalization stats

* Update for pre-commit

* Union compute_action_statistics code into update_action_statistics

* Update for pre-commit

* Add default value for dof_config and agent_pos_config

---------

Co-authored-by: farmer <farmer@x2robot.com>
This commit is contained in:
phoenixDong
2025-11-27 19:21:49 +08:00
committed by GitHub
co-authored by farmer
parent efc658f178
commit 908329d214
8 changed files with 155 additions and 69 deletions
Executable → Regular
View File
+100 -7
View File
@@ -5,14 +5,16 @@ This module provides utilities for preprocessing text, images, and actions
for multimodal transformer models in robotic learning tasks.
"""
import re
import torch
import random
from collections import OrderedDict
from typing import List, Dict, Any, Optional, Union, Tuple
from transformers import BatchFeature
from dataclasses import dataclass
import json
import random
import re
from collections import OrderedDict
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Tuple, Union
import numpy as np
import torch
from transformers import BatchFeature
KEY_MAPPINGS = {
"lerobot/aloha_mobile_cabinet": {
@@ -32,6 +34,15 @@ KEY_MAPPINGS = {
"state": "state",
"action": "actions",
},
"x2": {
"camera": {
"observation.images.faceImg": "face_view",
"observation.images.leftImg": "left_wrist_view",
"observation.images.rightImg": "right_wrist_view",
},
"state": "state",
"action": "actions",
},
}
CAMERA_NAME_MAPPING = {
@@ -662,3 +673,85 @@ def load_norm_stats(norm_stats_path, dataset_name):
)
return {"action": action_norm_stats, "state": state_norm_stats}
def update_action_statistics(
action_statistic_dof: Dict[str, Any],
norm_stats_path: str,
repo_id: str,
dof_config: Dict[str, int] = None,
agent_pos_config: Dict[str, int] = None,
robot_name: str = None,
customized_dof_config: Dict[str, int] = None,
customized_agent_pos_config: Dict[str, int] = None,
) -> None:
"""
Update the action statistics dictionary with new robot configuration.
Args:
action_statistic_dof (Dict[str, Any]): The dictionary to be updated with statistics
norm_stats_path (str): Path to the normalization statistics file
repo_id (str): Repository ID for the LeRobot configuration
dof_config (Dict[str, int]): Configuration mapping DOF names to their dimensions
agent_pos_config (Dict[str, int]): Configuration mapping agent position names to their dimensions
robot_name (str, optional): Name of the robot. If None, uses repo_id as the key
customized_dof_config (Dict[str, int], optional): Customized DOF configuration for specific robot
customized_agent_pos_config (Dict[str, int], optional): Customized agent position configuration for specific robot
"""
# Load normalization statistics
norm_stats = load_norm_stats(norm_stats_path, repo_id)
# Extract min and delta values for action and state
action_min = norm_stats["action"].min.numpy().tolist()
action_delta = norm_stats["action"].delta.numpy().tolist()
state_min = norm_stats["state"].min.numpy().tolist()
state_delta = norm_stats["state"].delta.numpy().tolist()
# Use customized configurations if provided, otherwise use default ones
current_dof_config = (
customized_dof_config if customized_dof_config is not None else dof_config
)
current_agent_pos_config = (
customized_agent_pos_config
if customized_agent_pos_config is not None
else agent_pos_config
)
# Prepare keys and values for DOF and agent position configurations
dof_key = []
agent_pos_key = []
dof_value = []
agent_pos_value = []
stats_dict = {}
# Extract DOF configuration
for k, v in current_dof_config.items():
dof_key.append(k)
dof_value.append(v)
# Extract agent position configuration
for k, v in current_agent_pos_config.items():
agent_pos_key.append(k)
agent_pos_value.append(v)
# Calculate DOF indices and extract corresponding min/delta values
dof_idx = np.array([0] + dof_value).cumsum()
for i in range(len(dof_idx) - 1):
stats_dict[dof_key[i]] = {
"min": action_min[dof_idx[i] : dof_idx[i + 1]],
"delta": action_delta[dof_idx[i] : dof_idx[i + 1]],
}
# Calculate agent position indices and extract corresponding min/delta values
agent_pos_idx = np.array([0] + agent_pos_value).cumsum()
for i in range(len(agent_pos_idx) - 1):
stats_dict[agent_pos_key[i]] = {
"min": state_min[agent_pos_idx[i] : agent_pos_idx[i + 1]],
"delta": state_delta[agent_pos_idx[i] : agent_pos_idx[i + 1]],
}
# Use provided robot name or repo_id as the key
robot_key = robot_name if robot_name is not None else repo_id
# Update the action_statistic_dof dictionary
action_statistic_dof.update({robot_key: stats_dict})
Executable → Regular
View File
@@ -44,8 +44,8 @@ from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl import (
Qwen2_5_VLSdpaAttention,
)
from wall_x.data.config import ACTION_DATASET_NAMES, MULTIMODAL_DATASET_NAMES
from wall_x.data.utils import update_action_statistics
from wall_x.utils.constant import action_statistic_dof
from wall_x.data.utils import load_norm_stats
from pprint import pprint
logger = logging.get_logger(__name__)
@@ -771,43 +771,19 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
"customized_agent_pos_config"
]
norm_stats_path = config["norm_stats_path"]
norm_stats = load_norm_stats(
norm_stats_path, config["data"]["lerobot_config"]["repo_id"]
)
action_min = norm_stats["action"].min.numpy().tolist()
action_delta = norm_stats["action"].delta.numpy().tolist()
state_min = norm_stats["state"].min.numpy().tolist()
state_delta = norm_stats["state"].delta.numpy().tolist()
# Use the compute_action_statistics function from utils
name = config["customized_robot_config"]["name"]
dof_key = []
agent_pos_key = []
dof_value = []
agent_pos_value = []
stats_dict = {}
for k, v in customized_dof_config.items():
dof_key.append(k)
dof_value.append(v)
for k, v in customized_agent_pos_config.items():
agent_pos_key.append(k)
agent_pos_value.append(v)
dof_idx = np.array([0] + dof_value).cumsum()
for i in range(len(dof_idx) - 1):
stats_dict[dof_key[i]] = {
"min": action_min[dof_idx[i] : dof_idx[i + 1]],
"delta": action_delta[dof_idx[i] : dof_idx[i + 1]],
}
agent_pos_idx = np.array([0] + agent_pos_value).cumsum()
for i in range(len(agent_pos_idx) - 1):
stats_dict[agent_pos_key[i]] = {
"min": state_min[agent_pos_idx[i] : agent_pos_idx[i + 1]],
"delta": state_delta[agent_pos_idx[i] : agent_pos_idx[i + 1]],
}
action_statistic_dof[name] = stats_dict
update_action_statistics(
action_statistic_dof=action_statistic_dof, # Assuming this is a global variable
norm_stats_path=norm_stats_path,
repo_id=config["data"]["lerobot_config"]["repo_id"],
robot_name=name,
customized_dof_config=customized_dof_config,
customized_agent_pos_config=customized_agent_pos_config,
)
print("Customized robot config added")
pprint(action_statistic_dof)
+40 -23
View File
@@ -16,12 +16,9 @@ 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.data.utils import update_action_statistics
from wall_x.utils.constant import action_statistic_dof
from wall_x.model.action_head import Normalizer
try:
import msgpack
@@ -45,17 +42,25 @@ 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"):
def __init__(
self,
config_path: str,
uri: str = "ws://localhost:8000",
norm_stats_path: str = "x2_norm_stats.json",
):
"""Initialize client.
Args:
config_path: Path to train config file
uri: WebSocket URI of the server (e.g., ws://localhost:8000)
norm_stats_path: Path to normalization stats file
"""
self.uri = uri
self.websocket = None
self.metadata = None
self._loop = None
self._thread = None
self.norm_stats_path = norm_stats_path
with open(config_path, "r") as f:
self.train_config = yaml.load(f, Loader=yaml.FullLoader)
@@ -166,20 +171,21 @@ class WallXClient:
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)
# Define default configurations
dof_config = {"biarm_eed_with_base": 20}
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")
agent_pos_config = {"biarm_eed_with_base": 20}
update_action_statistics(
action_statistic_dof=action_statistic_dof,
norm_stats_path=self.norm_stats_path,
repo_id="x2",
dof_config=dof_config,
agent_pos_config=agent_pos_config,
)
self.normalizer_action = Normalizer(action_statistic_dof, dof_config)
self.normalizer_propri = Normalizer(action_statistic_dof, agent_pos_config)
print("Normalizer initialized")
@@ -213,6 +219,8 @@ def prepare_batch_sync(data, normalizer_action, normalizer_propri, dataset_names
def init_serving_sample_dataset(train_config):
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
repo_id = train_config["data"]["lerobot_config"]["repo_id"]
meta_info = LeRobotDatasetMetadata(repo_id)
@@ -237,7 +245,9 @@ def main_sync(args):
"""Synchronous version of main function."""
# Create client and connect
client = WallXClient(args.config_path, uri=args.uri)
client = WallXClient(
args.config_path, uri=args.uri, norm_stats_path=args.norm_stats_path
)
client.connect_sync()
dataset, repo_id = init_serving_sample_dataset(client.train_config)
@@ -295,7 +305,9 @@ def main_sync(args):
async def main(args):
client = WallXClient(args.config_path, uri=args.uri)
client = WallXClient(
args.config_path, uri=args.uri, norm_stats_path=args.norm_stats_path
)
await client.connect()
dataset, repo_id = init_serving_sample_dataset(client.train_config)
@@ -363,14 +375,19 @@ if __name__ == "__main__":
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",
default="config_from_qwen_libero.yml",
help="Train config path",
)
parser.add_argument(
"--save_dir",
default="/x2robot_v2/vincent/workspace/opensource/plots/libero",
default="libero",
help="Save directory",
)
parser.add_argument(
"--norm_stats_path",
default="x2_norm_stats.json",
help="Normalization stats path",
)
args = parser.parse_args()
# Synchronous mode
+1 -1
View File
@@ -30,7 +30,7 @@ class WallXPolicy(BasePolicy):
min_pixels: int = 4 * 28 * 28,
max_pixels: int = 16384 * 28 * 28,
image_factor: int = 28,
max_length: int = 768,
max_length: int = 2048,
):
"""Initialize the Wall-X policy.
+3 -3
View File
@@ -3,7 +3,7 @@ import http
import logging
import time
import traceback
from typing import Dict, Any
from typing import Any, Dict, Optional
try:
import msgpack
@@ -54,7 +54,7 @@ class WebsocketPolicyServer:
policy: BasePolicy,
host: str = "0.0.0.0",
port: int = 8000,
metadata: Dict | None = None,
metadata: Optional[Dict] = None,
) -> None:
self._policy = policy
self._host = host
@@ -126,7 +126,7 @@ class WebsocketPolicyServer:
def _health_check(
connection: _server.ServerConnection, request: _server.Request
) -> _server.Response | None:
) -> Optional[_server.Response]:
if request.path == "/healthz":
return connection.respond(http.HTTPStatus.OK, "OK\n")
return None
Executable → Regular
View File