Files
Mujoco_WASM/training_server/mobile_manipulator/train.py
T
chenlin f3a8a38acd
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
2026-09-24 09:57:41 +08:00

180 lines
7.3 KiB
Python

"""Server-owned SB3 PPO runner. stdout follows the shared iteration/scalar protocol."""
import argparse
import json
from pathlib import Path
import numpy as np
import torch
from .bootstrap import bootstrap_navigation
from .config import MOBILE_TASKS, validate_mobile_params
from .env import MobileManipulatorEnv
from .evaluation import evaluate_policy
from .export_onnx import export_policy
class DeterministicActor(torch.nn.Module):
def __init__(self, policy):
super().__init__()
self.policy = policy
def forward(self, observation):
return self.policy(observation, deterministic=True)[0]
def main():
from stable_baselines3 import PPO
from stable_baselines3.common.monitor import Monitor
from stable_baselines3.common.vec_env import DummyVecEnv
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--package", required=True)
parser.add_argument("--iterations", type=int, default=1000)
parser.add_argument("--num-envs", type=int, default=1)
parser.add_argument("--output", required=True)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--device", default="cpu")
parser.add_argument("--task-id", required=True, choices=MOBILE_TASKS)
parser.add_argument("--params", default="{}")
parser.add_argument("--resume", help="server-owned trusted PPO checkpoint only")
args = parser.parse_args()
torch.set_num_threads(min(4, torch.get_num_threads()))
params = validate_mobile_params(json.loads(args.params))
if args.iterations < 1 or not 1 <= args.num_envs <= 64:
parser.error("iterations/num-envs out of range")
if args.device.startswith("cuda") and not torch.cuda.is_available():
raise RuntimeError("选择了 GPU,但训练 Python 中 CUDA 不可用;请改用 CPU")
if params["stage"] != "navigate" and not args.resume:
parser.error("later stages require a validated previous-stage checkpoint")
reset_options = {"object": params["objectPosition"], "goal": params["goalPosition"]}
print(f"Training stage: {params['stage']} | rate-limited-position-target-v1", flush=True)
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
class ConsolePPO(PPO):
iteration = 0
def train(self):
super().train()
self.iteration += 1
print(f"Learning iteration {self.iteration} / {args.iterations}", flush=True)
print(f"Total timesteps: {self.num_timesteps}", flush=True)
for key, label in (
("value_loss", "value loss"),
("policy_gradient_loss", "surrogate loss"),
("entropy_loss", "entropy loss"),
):
value = self.logger.name_to_value.get(f"train/{key}")
if value is not None and np.isfinite(value):
print(f"Mean {label}: {value:.9g}", flush=True)
print(f"Mean reward: {np.mean(self.rollout_buffer.rewards):.9g}", flush=True)
if self.ep_info_buffer:
success = np.mean([x["is_success"] for x in self.ep_info_buffer])
peak = max(x["max_joint_velocity"] for x in self.ep_info_buffer)
print(f"Mean success rate: {success:.6g}", flush=True)
print(f"Max joint velocity: {peak:.6g}", flush=True)
print(
f"Mean episode length: {np.mean([x['l'] for x in self.ep_info_buffer]):.9g}",
flush=True,
)
def make_env():
return Monitor(
MobileManipulatorEnv(
args.package,
reset_options=reset_options,
stage=params["stage"],
position_jitter=params["positionJitter"],
),
info_keywords=("is_success", "max_joint_velocity", "safety_stop"),
)
# Native MuJoCo CPU simulation, vectorized rollout; device selects the PPO network.
env = DummyVecEnv([make_env for _ in range(args.num_envs)])
try:
if env.envs[0].unwrapped.config["id"] != MOBILE_TASKS[args.task_id]:
raise ValueError("task/robot mismatch")
rollout = params["rolloutSteps"] * args.num_envs
batch = min(64, rollout)
while rollout % batch:
batch -= 1
settings = dict(
seed=args.seed,
verbose=0,
device=args.device,
n_steps=params["rolloutSteps"],
batch_size=batch,
)
if args.resume:
agent = ConsolePPO.load(args.resume, env=env, **settings)
agent.iteration = 0
if agent.ep_info_buffer is not None:
agent.ep_info_buffer.clear()
if getattr(agent, "training_stage", params["stage"]) != params["stage"]:
# Explore the newly enabled arm without destroying base navigation.
with torch.no_grad():
agent.policy.log_std[3:].fill_(-1.5)
else:
agent = ConsolePPO(
"MlpPolicy",
env,
**settings,
policy_kwargs=dict(
net_arch=[128, 128],
activation_fn=torch.nn.ELU,
log_std_init=-3.0 if params["navigationBootstrapSteps"] else -1.0,
),
# Conservative PPO updates: do not destroy navigation initialization
# while the critic is still learning the sparse terminal return.
learning_rate=3e-5,
n_epochs=4,
ent_coef=0.001,
target_kl=0.005,
)
if not args.resume:
bootstrap_navigation(agent, args.package, params, reset_options, args.seed + 50_000)
agent.learn(
total_timesteps=args.iterations * rollout,
log_interval=None,
reset_num_timesteps=not bool(args.resume),
)
agent.training_stage = params["stage"]
agent.save(str(output.with_suffix(".ppo.zip")))
# Distinct seeds, deterministic actions and the same controller used by deployment.
evaluation = evaluate_policy(
agent, args.package, params, reset_options, args.seed + 100_000
)
(output.parent / "evaluation.json").write_text(json.dumps(evaluation, indent=2) + "\n")
metadata = export_policy(
DeterministicActor(agent.policy), args.package, output, params["stage"]
)
metadata.update(
version=1,
browserCompatible=True,
trainingTaskId=args.task_id,
seed=args.seed,
resetOptions=reset_options,
trainingParams=params,
trainedTimesteps=agent.num_timesteps,
initialization="resumed-checkpoint"
if args.resume
else "navigation-BC-then-PPO"
if params["navigationBootstrapSteps"]
else "PPO-from-scratch",
evaluation={k: v for k, v in evaluation.items() if k != "rollouts"},
)
output.with_name("deployment.json").write_text(json.dumps(metadata, indent=2) + "\n")
print(
f"Evaluation success rate: {evaluation['successRate']:.3f}; "
f"safety stops: {evaluation['safetyStops']}",
flush=True,
)
print("ONNX export complete (export does NOT imply task competence)", flush=True)
finally:
env.close()
if __name__ == "__main__":
main()