Merge branch 'lab'

This commit is contained in:
wty-yy
2026-05-18 22:36:32 +08:00
6 changed files with 140 additions and 10 deletions

View File

@@ -1,4 +1,8 @@
# UPDATE # UPDATE
## 20260518
### v1.1.7
1. 新增go2 lab环境评估
2. 新增server, client通讯检测
## 20260324 ## 20260324
### v1.1.6 ### v1.1.6
1. 修复multi_pipeline中统计terrain_weighted_summary的地形等级加权分错误但不影响stress_pipeline最终计算的benchmark_score只影响metrics里面的mean和mean@50的计算 1. 修复multi_pipeline中统计terrain_weighted_summary的地形等级加权分错误但不影响stress_pipeline最终计算的benchmark_score只影响metrics里面的mean和mean@50的计算

View File

@@ -6,12 +6,36 @@ from typing import Dict, Any, Optional
from robogauge.scripts.server import ResponseStatus from robogauge.scripts.server import ResponseStatus
class RoboGaugeClient: class RoboGaugeClient:
def __init__(self, base_url: str = "http://127.0.0.1:9973"): def __init__(self, base_url: str = "http://127.0.0.1:9973", request_timeout: float = 5.0):
self.base_url = base_url self.base_url = base_url
self.request_timeout = request_timeout
self.processing_ids = [] self.processing_ids = []
self.task_id2info = {} self.task_id2info = {}
self.response_data = {} self.response_data = {}
def is_server_available(self) -> bool:
"""Return True when the RoboGauge API and evaluator main process are healthy."""
try:
response = requests.get(
f"{self.base_url}/health",
timeout=self.request_timeout,
)
if response.status_code != 200:
return False
return response.json().get("status") == "ok"
except (requests.exceptions.RequestException, ValueError):
return False
def wait_until_available(self, retry_interval: int = 2) -> None:
"""Block until the RoboGauge API can be reached."""
while not self.is_server_available():
print(
f"[RoboGaugeClient]⏳ Waiting for RoboGauge server at {self.base_url} "
f"(retrying in {retry_interval}s)..."
)
time.sleep(retry_interval)
print(f"[RoboGaugeClient]✅ Connected to RoboGauge server at {self.base_url}.")
def submit_task(self, def submit_task(self,
model_path: str, model_path: str,
step: int, step: int,
@@ -40,9 +64,16 @@ class RoboGaugeClient:
print(f"[RoboGaugeClient]📤 Preparing to submit task: {task_name}") print(f"[RoboGaugeClient]📤 Preparing to submit task: {task_name}")
if wait_for_server:
self.wait_until_available(retry_interval=retry_interval)
while True: while True:
try: try:
response = requests.post(f"{self.base_url}/submit_eval", json=payload) response = requests.post(
f"{self.base_url}/submit_eval",
json=payload,
timeout=self.request_timeout,
)
if response.status_code == 200: if response.status_code == 200:
data = response.json() data = response.json()
@@ -55,9 +86,9 @@ class RoboGaugeClient:
print(f"[RoboGaugeClient]❌ Server returned error: {response.text}") print(f"[RoboGaugeClient]❌ Server returned error: {response.text}")
return None return None
except requests.exceptions.ConnectionError: except requests.exceptions.RequestException:
if not wait_for_server: if not wait_for_server:
print("[RoboGaugeClient]❌ Unable to connect to server (Connection Refused).") print("[RoboGaugeClient]❌ Unable to connect to server.")
return None return None
print(f"[RoboGaugeClient]⏳ Server not responding, retrying in {retry_interval} seconds... (Make sure server.py is running)") print(f"[RoboGaugeClient]⏳ Server not responding, retrying in {retry_interval} seconds... (Make sure server.py is running)")
@@ -66,11 +97,24 @@ class RoboGaugeClient:
print(f"[RoboGaugeClient]❌ Unknown error: {e}") print(f"[RoboGaugeClient]❌ Unknown error: {e}")
return None return None
def monitor_tasks(self): def monitor_tasks(self, wait_for_server: bool = True, retry_interval: int = 2):
print("[RoboGaugeClient]⏱️ Monitoring submitted tasks...") print("[RoboGaugeClient]⏱️ Monitoring submitted tasks...")
""" Monitor all submitted tasks until completion. """ """ Monitor all submitted tasks until completion. """
if wait_for_server:
self.wait_until_available(retry_interval=retry_interval)
for task_id in reversed(self.processing_ids): for task_id in reversed(self.processing_ids):
respone = requests.get(f"{self.base_url}/get_result/{task_id}") try:
respone = requests.get(f"{self.base_url}/get_result/{task_id}", timeout=self.request_timeout)
except requests.exceptions.RequestException:
if not wait_for_server:
raise
print(
f"[RoboGaugeClient]⏳ Lost connection to RoboGauge server at {self.base_url}; "
f"waiting before checking task results..."
)
self.wait_until_available(retry_interval=retry_interval)
continue
if respone.status_code != 200: if respone.status_code != 200:
continue continue
resp_data = respone.json() resp_data = respone.json()

View File

@@ -21,11 +21,12 @@ os.environ["OMP_NUM_THREADS"] = "1"
os.environ["MKL_NUM_THREADS"] = "1" os.environ["MKL_NUM_THREADS"] = "1"
import multiprocessing import multiprocessing
import threading
import uvicorn import uvicorn
import queue import queue
import time import time
import uuid import uuid
from fastapi import FastAPI from fastapi import FastAPI, HTTPException
from pydantic import BaseModel from pydantic import BaseModel
from typing import Dict, Optional from typing import Dict, Optional
import argparse import argparse
@@ -47,6 +48,8 @@ default_args_list = [
'--headless', '--headless',
] ]
HEARTBEAT_TIMEOUT_SECONDS = 10.0
@dataclass @dataclass
class EvalTaskData: class EvalTaskData:
model_path: str model_path: str
@@ -67,15 +70,42 @@ class ResponseStatus:
ERROR = "error" ERROR = "error"
NOT_FOUND = "not_found" NOT_FOUND = "not_found"
def run_api_server(input_queue: multiprocessing.Queue, result_dict: dict, port=9973): def update_main_heartbeat(health_dict: dict, stop_event: threading.Event, interval: float = 1.0):
while not stop_event.is_set():
health_dict["main_pid"] = os.getpid()
health_dict["last_heartbeat"] = time.time()
stop_event.wait(interval)
def is_main_process_healthy(health_dict: dict) -> bool:
last_heartbeat = float(health_dict.get("last_heartbeat", 0.0))
return time.time() - last_heartbeat <= HEARTBEAT_TIMEOUT_SECONDS
def run_api_server(input_queue: multiprocessing.Queue, result_dict: dict, health_dict: dict, port=9973):
""" """
Running in a separate subprocess. Running in a separate subprocess.
I/O Process: submit requests -> put into queue -> return ID. I/O Process: submit requests -> put into queue -> return ID.
""" """
app = FastAPI() app = FastAPI()
@app.get("/health")
def health():
last_heartbeat = float(health_dict.get("last_heartbeat", 0.0))
heartbeat_age = time.time() - last_heartbeat
if heartbeat_age > HEARTBEAT_TIMEOUT_SECONDS:
raise HTTPException(
status_code=503,
detail=f"RoboGauge main process heartbeat is stale ({heartbeat_age:.1f}s).",
)
return {
"status": "ok",
"main_pid": health_dict.get("main_pid"),
"heartbeat_age": heartbeat_age,
}
@app.post("/submit_eval") @app.post("/submit_eval")
def submit_eval(req: EvalRequest): def submit_eval(req: EvalRequest):
if not is_main_process_healthy(health_dict):
raise HTTPException(status_code=503, detail="RoboGauge main process is not healthy.")
task_id = str(uuid.uuid4()) task_id = str(uuid.uuid4())
task_data = EvalTaskData( task_data = EvalTaskData(
model_path=req.model_path, model_path=req.model_path,
@@ -110,10 +140,18 @@ def main():
manager = ctx.Manager() manager = ctx.Manager()
task_queue = manager.Queue() task_queue = manager.Queue()
results_store = manager.dict() results_store = manager.dict()
health_store = manager.dict()
heartbeat_stop = threading.Event()
heartbeat_thread = threading.Thread(
target=update_main_heartbeat,
args=(health_store, heartbeat_stop),
daemon=True,
)
heartbeat_thread.start()
api_p = ctx.Process( api_p = ctx.Process(
target=run_api_server, target=run_api_server,
args=(task_queue, results_store, args_cli.port), args=(task_queue, results_store, health_store, args_cli.port),
daemon=True daemon=True
) )
api_p.start() api_p.start()
@@ -165,6 +203,8 @@ def main():
print("\n🛑 Shutting down...") print("\n🛑 Shutting down...")
api_p.terminate() api_p.terminate()
api_p.join() api_p.join()
heartbeat_stop.set()
heartbeat_thread.join(timeout=2.0)
if __name__ == "__main__": if __name__ == "__main__":
main() main()

View File

@@ -1,6 +1,14 @@
from robogauge.utils.task_register import task_register from robogauge.utils.task_register import task_register
from robogauge.tasks.simulator.mujoco_config import MujocoConfig from robogauge.tasks.simulator.mujoco_config import MujocoConfig
from robogauge.tasks.robots import RobotConfig, Go2Config, Go2MoEConfig, Go2TerrainConfig, Go2MoETerrainConfig from robogauge.tasks.robots import (
RobotConfig,
Go2Config,
Go2LabConfig,
Go2MoEConfig,
Go2TerrainConfig,
Go2LabTerrainConfig,
Go2MoETerrainConfig,
)
from robogauge.tasks.pipeline import BasePipeline from robogauge.tasks.pipeline import BasePipeline
from robogauge.tasks.gauge import BaseGaugeConfig from robogauge.tasks.gauge import BaseGaugeConfig
@@ -26,3 +34,12 @@ task_register.register('go2_moe.wave', BasePipeline, MujocoConfig, Go2WaveGaugeC
task_register.register('go2_moe.stairs_fd', BasePipeline, MujocoConfig, Go2StairsForwardGaugeConfig, Go2MoETerrainConfig) task_register.register('go2_moe.stairs_fd', BasePipeline, MujocoConfig, Go2StairsForwardGaugeConfig, Go2MoETerrainConfig)
task_register.register('go2_moe.stairs_bd', BasePipeline, MujocoConfig, Go2StairsBackwardGaugeConfig, Go2MoETerrainConfig) task_register.register('go2_moe.stairs_bd', BasePipeline, MujocoConfig, Go2StairsBackwardGaugeConfig, Go2MoETerrainConfig)
task_register.register('go2_moe.obstacle', BasePipeline, MujocoConfig, Go2ObstacleGaugeConfig, Go2MoETerrainConfig) task_register.register('go2_moe.obstacle', BasePipeline, MujocoConfig, Go2ObstacleGaugeConfig, Go2MoETerrainConfig)
# Go2 Lab
task_register.register('go2_lab.flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2LabConfig)
task_register.register('go2_lab.slope_fd', BasePipeline, MujocoConfig, Go2SlopeForwardGaugeConfig, Go2LabTerrainConfig)
task_register.register('go2_lab.slope_bd', BasePipeline, MujocoConfig, Go2SlopeBackwardGaugeConfig, Go2LabTerrainConfig)
task_register.register('go2_lab.wave', BasePipeline, MujocoConfig, Go2WaveGaugeConfig, Go2LabTerrainConfig)
task_register.register('go2_lab.stairs_fd', BasePipeline, MujocoConfig, Go2StairsForwardGaugeConfig, Go2LabTerrainConfig)
task_register.register('go2_lab.stairs_bd', BasePipeline, MujocoConfig, Go2StairsBackwardGaugeConfig, Go2LabTerrainConfig)
task_register.register('go2_lab.obstacle', BasePipeline, MujocoConfig, Go2ObstacleGaugeConfig, Go2LabTerrainConfig)

View File

@@ -1,6 +1,7 @@
from .base_robot_config import RobotConfig from .base_robot_config import RobotConfig
from .base_robot import BaseRobot from .base_robot import BaseRobot
from .go2.go2_config import Go2Config, Go2TerrainConfig from .go2.go2_config import Go2Config, Go2TerrainConfig
from .go2.go2_lab_config import Go2LabConfig, Go2LabTerrainConfig
from .go2.go2 import Go2 from .go2.go2 import Go2
from .go2.go2_moe_config import Go2MoEConfig, Go2MoETerrainConfig from .go2.go2_moe_config import Go2MoEConfig, Go2MoETerrainConfig
from .go2.go2_moe import Go2MoE from .go2.go2_moe import Go2MoE

View File

@@ -0,0 +1,24 @@
# -*- coding: utf-8 -*-
'''
@File : go2_lab_config.py
@Time : 2026/04/11 17:59:09
@Author : wty-yy
@Version : 1.0
@Blog : https://wty-yy.github.io/
@Desc : Go2 Lab Robot Configuration
'''
from robogauge.tasks.robots.go2.go2_config import Go2Config, Go2TerrainConfig
class Go2LabConfig(Go2Config):
"""Go2 Lab robot configuration aligned with RobotLab observation scaling."""
class control(Go2Config.control):
p_gains = [25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0, 25.0] # [N*m/rad]
class scales(Go2Config.control.scales):
# RobotLab policy command observation uses unit scale.
cmd = [1.0, 1.0, 1.0]
class Go2LabTerrainConfig(Go2LabConfig, Go2TerrainConfig):
"""Go2 Lab robot configuration for terrain tasks."""