From 2aa39ff5f8ef303a9ae5d3c696f57aa5f61245c2 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Tue, 29 Sep 2026 15:32:49 -0400 Subject: [PATCH] docs: the algorithm, policy and contributing pages match the code - six algorithms, the frame rule, resuming, and the full algorithm and policy contracts --- docs/algorithm/custom_algorithm.md | 109 +++++++++++++++++----- docs/algorithm/custom_algorithm.zh.md | 103 ++++++++++++++------ docs/algorithm/index.md | 66 ++++++++++--- docs/algorithm/index.zh.md | 59 +++++++++--- docs/algorithm/ppo.md | 129 +++++++++++++++++++++----- docs/algorithm/ppo.zh.md | 114 ++++++++++++++++++----- docs/contributing/algorithm.md | 11 ++- docs/contributing/algorithm.zh.md | 8 +- docs/contributing/policy.md | 5 +- docs/contributing/policy.zh.md | 4 +- docs/policy/custom_policy.md | 22 +++++ docs/policy/custom_policy.zh.md | 22 ++++- docs/policy/dppo_policy.md | 114 +++++++++++++++++++---- docs/policy/dppo_policy.zh.md | 104 +++++++++++++++++---- docs/policy/index.md | 36 +++++-- docs/policy/index.zh.md | 33 +++++-- 16 files changed, 758 insertions(+), 181 deletions(-) diff --git a/docs/algorithm/custom_algorithm.md b/docs/algorithm/custom_algorithm.md index 7a5089b..7d6c714 100644 --- a/docs/algorithm/custom_algorithm.md +++ b/docs/algorithm/custom_algorithm.md @@ -6,9 +6,11 @@ Add a server-side algorithm so `plugrl-run-server` can select it by UID. 1. Create a package under `plugrl-server/src/plugrl_server/algorithm//`. 2. Register a config dataclass and an algorithm class. -3. Ensure the module is imported at server startup so Tyro can discover it. +3. Ensure both modules are imported at server startup so Tyro can discover them. -Reference implementation: `plugrl-server/examples/sac/sac.py`. +Reference implementation: `plugrl-server/examples/sac/sac.py`, an off-policy +SAC with its own policy in `sac_policy.py`. Run it from `plugrl-server` with +`python examples/sac/sac.py sac_policy default sac default`. ## File layout @@ -16,37 +18,79 @@ Put the code in one of these layouts. ### Built-in in `plugrl-server` -- `plugrl_server/algorithm//.py` implementation -- `plugrl_server/algorithm//__init__.py` registration import -- `plugrl_server/algorithm/__init__.py` imports your package +- `plugrl_server/algorithm//_config.py`: the config + dataclass, registered with `@register_algo_config` +- `plugrl_server/algorithm//.py`: the algorithm class, + registered with `@register_algo` +- `plugrl_server/algorithm/__init__.py`: imports both modules + +Import both. The config module alone puts the UID in the CLI, but the run +prints its config and dies with `KeyError: 'Algorithm is not registered.'`. +One module holding both decorators, as `examples/sac/sac.py` does, is also fine. ### Plug-in in your own package - Put the algorithm module in your own Python package. -- Import the module before calling `plugrl_server.cli:main`. +- Import it, config and class, before calling `plugrl_server.cli:main`. ## Server contract The WebSocket server loop calls these methods. -- `infer(obs) -> (action, runtime_state)` -- `feedback(...) -> (prev_node, global_step, log_dict)` -- `learn() -> (global_step, log_dict)` -- Scheduling and checkpoint hooks: `should_learn`, `should_save`, `should_stop`, `create_checkpoint`, `load_checkpoint` +- `infer(obs) -> (action, runtime_state)`, once per batch of environments. +- `derive_train_state(runtime_state) -> train_state`, right after every + `infer`. The default returns `None`. The result is sliced per environment + and comes back as `feedback(train_state=...)`. `example_train_state(n)` + calls it on `policy.fake_runtime_state(n)`, which is how the built-in + algorithms size their buffers. +- `feedback(...) -> (prev_node, global_step, log_dict)`, only for frames that + will be stored. +- `discard_feedback(*, info, next_terminated, next_truncated)`, for every + other frame. The default records the finished episode's metrics and nothing + else. If you override it, call `super().discard_feedback(...)`, or those + episodes are missing from `rollout/*`. +- `learn() -> (global_step, log_dict)`, with `pre_learn()` before it and + `post_learn()` after it. The base `post_learn` resets the episode metrics, + so an override should call `super().post_learn()`. +- Scheduling and checkpoint hooks: `should_learn`, `should_save`, `should_stop`, `create_checkpoint` + +The CLI calls two more before the server starts. + +- `init_optimizers()`, right after the algorithm is built. The default does + nothing. `dppo` builds its optimizers here. +- `load_checkpoint(checkpoint)`, when the run is started with `--resume`. + +And one class attribute. + +- `on_policy: bool = True`. The server passes a frame to `feedback` only if + its action came from the current policy and `should_learn()` is false. + Everything else goes to `discard_feedback`. An algorithm that learns from a + replay buffer sets `on_policy = False` and is given every frame that has a + step state, as the SAC example does. See + [Training loop](ppo.md#what-happens-on-the-server). + +What the loop implies for the hooks. + +- The server does not infer while `should_learn()` is true. If it is still + true after `learn()`, the server learns again without inferring in between. + So `learn()` or `post_learn()` must empty the buffer, or move whatever + counter `should_learn` reads. +- `should_save()` and `should_stop()` are checked on every pass of the loop, + not once per learn step. `should_save()` must turn false once + `create_checkpoint()` has run. The built-ins record the saved iteration + inside `create_checkpoint`. See `plugrl_server/algorithm/base_algorithm.py` for exact signatures. The -server does call `learn()`, but `learn()` is concrete on `BaseAlgorithm`: it -calls `learn_impl()` and then wraps the result with `build_train_info`. -`learn_impl` is the abstract method, so that is the one you override. -Overriding `learn` instead leaves `learn_impl` unimplemented and the class -abstract, and `make_algo` fails with `TypeError`. An earlier version of the -template below did exactly that. +server calls `learn()`, but `learn()` is concrete on `BaseAlgorithm`: it calls +`learn_impl()` and then wraps the result with `build_train_info`, which adds +the `rollout/*` metrics. `learn_impl` is the abstract method, so that is the +one you override. Overriding `learn` instead leaves `learn_impl` +unimplemented and the class abstract, and `make_algo` fails with `TypeError`. `infer` and `feedback` take and return `PolicyRuntimeState` from `plugrl_server.policy.state`; `feedback` also takes `train_state: -PolicyTrainState = None`. There is no `InternalState` type and no -`get_action_and_internal_state` method anywhere in `plugrl-server` - the -template used both names and neither imports. Every `feedback` parameter is +PolicyTrainState = None`. Older templates used `InternalState` and +`get_action_and_internal_state`; neither exists. Every `feedback` parameter is keyword-only, so a mismatched name is a `TypeError` on the server's first call, not a silent rename. @@ -60,7 +104,7 @@ import numpy as np from plugrl_server.algorithm.base_algorithm import BaseAlgoConfig, BaseAlgorithm from plugrl_server.algorithm.registration import register_algo, register_algo_config from plugrl_server.common.checkpoint_manager import Checkpoint -from plugrl_server.policy.base_policy import BasePolicy +from plugrl_server.policy.base_torch_policy import BaseTorchPolicy from plugrl_server.policy.state import PolicyRuntimeState, PolicyTrainState UID = "your-algo" @@ -69,12 +113,16 @@ UID = "your-algo" @register_algo_config(UID) @dataclasses.dataclass class YourAlgoConfig(BaseAlgoConfig): - total_timesteps: int = 100_000 + # The run length is BaseAlgoConfig's field, set with --algo.global-steps. + global_steps: int | None = 100_000 @register_algo(UID) class YourAlgorithm(BaseAlgorithm): - def __init__(self, config: YourAlgoConfig, policy: BasePolicy): + # Set to False if the algorithm learns from a replay buffer. + on_policy = True + + def __init__(self, config: YourAlgoConfig, policy: BaseTorchPolicy): super().__init__(config=config, policy=policy) self.global_step = 0 @@ -97,6 +145,11 @@ class YourAlgorithm(BaseAlgorithm): next_truncated: bool, prev_node: tuple, ) -> tuple[tuple, int, dict]: + # Store the frame here. + if (next_terminated or next_truncated) and "episode" in info: + if bool(info["episode"].get("mask", True)): + # Without this, rollout/success, reward and length stay 0. + self.record_episode_metrics(info["episode"]) self.global_step += 1 return prev_node, self.global_step, {} @@ -107,12 +160,13 @@ class YourAlgorithm(BaseAlgorithm): return False def should_stop(self) -> bool: - return self.global_step >= self.config.total_timesteps + return self.global_step >= self.config.global_steps def should_save(self) -> bool: return False def create_checkpoint(self) -> Checkpoint: + # state_dict() exists because a BaseTorchPolicy is a torch.nn.Module. return Checkpoint(step=self.global_step, model=self.policy.state_dict(), optimizer=None, meta={}) def load_checkpoint(self, checkpoint: Checkpoint) -> None: @@ -124,6 +178,7 @@ class YourAlgorithm(BaseAlgorithm): ## Design rules - Put model and action generation parameters in the policy config. +- Use `BaseAlgoConfig.global_steps` for the run length rather than a field of your own. The progress display reads it. - Keep two counters if your `learn()` uses external data: environment steps and update steps. - Save schedule counters in `Checkpoint.meta` and restore them in `load_checkpoint`. - Use `BaseAlgorithm` unless you implement the distributed hooks required by `DDPAlgorithm`. @@ -134,7 +189,7 @@ Start with a smoke test. ```bash plugrl-run-server dummy-policy default your-algo default -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` For plug-in algorithms, import before entering the CLI. @@ -146,7 +201,11 @@ python -c "import my_pkg.plugrl_algorithms; from plugrl_server.cli import main; ## Troubleshooting -- Algorithm UID not listed in `plugrl-run-server --help`: module import did not run. +- Algorithm UID not listed in `plugrl-run-server --help`: module import did not run. +- `KeyError: 'Algorithm your-algo is not registered.'` after the config is printed: the class module was not imported, only the config module. +- `rollout/*` stays at 0: `feedback` does not call `record_episode_metrics`. If only some episodes are missing, an override of `discard_feedback` does not call `super()`. +- The server learns again and again without inferring: `should_learn()` is still true after `learn()`. +- A checkpoint is written on every loop pass: `should_save()` stays true after `create_checkpoint()`. - Duplicate flags under `--algo.*` and `--policy.*`: keep the parameter in one config. - After resume, learning or saving cadence drifts: restore all counters from `Checkpoint.meta`. - `DDPAlgorithm` errors at runtime: switch to `BaseAlgorithm` or implement the required distributed hooks. diff --git a/docs/algorithm/custom_algorithm.zh.md b/docs/algorithm/custom_algorithm.zh.md index db89220..6b0f92d 100644 --- a/docs/algorithm/custom_algorithm.zh.md +++ b/docs/algorithm/custom_algorithm.zh.md @@ -6,46 +6,82 @@ 1. 在 `plugrl-server/src/plugrl_server/algorithm//` 下新建包。 2. 注册一个配置 dataclass 和一个算法类。 -3. 确保 server 启动时会 import 该模块,Tyro 才能发现它。 +3. 确保 server 启动时这两个模块都会被 import,Tyro 才能发现它们。 -参考实现:`plugrl-server/examples/sac/sac.py`。 +参考实现:`plugrl-server/examples/sac/sac.py`,一个 off-policy 的 SAC,配套策略在 +`sac_policy.py`。在 `plugrl-server` 目录下用 +`python examples/sac/sac.py sac_policy default sac default` 运行。 -## 文件结构 +## 文件结构 {#file-layout} 下面两种组织方式都可以。 ### 直接放进 `plugrl-server` -- `plugrl_server/algorithm//.py` 实现 -- `plugrl_server/algorithm//__init__.py` 注册导入 -- `plugrl_server/algorithm/__init__.py` 导入你的包 +- `plugrl_server/algorithm//_config.py`:配置 dataclass,用 + `@register_algo_config` 注册 +- `plugrl_server/algorithm//.py`:算法类,用 `@register_algo` 注册 +- `plugrl_server/algorithm/__init__.py`:把这两个模块都 import 进来 + +两个都要 import。只 import 配置模块的话,UID 会出现在 CLI 里,但运行时打印完配置 +就报 `KeyError: 'Algorithm is not registered.'`。像 +`examples/sac/sac.py` 那样把两个装饰器放在同一个模块里也可以。 ### 放在你自己的包里 - 算法代码放进你自己的 Python 包。 -- 进入 `plugrl_server.cli:main` 之前先 import 一次。 +- 进入 `plugrl_server.cli:main` 之前先 import,配置和类都要。 ## Server 调用契约 WebSocket server loop 会调用这些方法。 -- `infer(obs) -> (action, runtime_state)` -- `feedback(...) -> (prev_node, global_step, log_dict)` -- `learn() -> (global_step, log_dict)` -- 调度与保存:`should_learn`、`should_save`、`should_stop`、`create_checkpoint`、`load_checkpoint` - -准确签名见 `plugrl_server/algorithm/base_algorithm.py`。server 确实调用的是 +- `infer(obs) -> (action, runtime_state)`:每批环境调用一次。 +- `derive_train_state(runtime_state) -> train_state`:每次 `infer` 之后紧接着调用, + 默认返回 `None`。结果按环境切开,再作为 `feedback(train_state=...)` 传回来。 + `example_train_state(n)` 会拿 `policy.fake_runtime_state(n)` 调它,内置算法就是 + 这样确定 buffer 大小的。 +- `feedback(...) -> (prev_node, global_step, log_dict)`:只对要存下来的帧调用。 +- `discard_feedback(*, info, next_terminated, next_truncated)`:其余每一帧都走这里。 + 默认实现只记录已结束 episode 的指标。覆写时要调用 + `super().discard_feedback(...)`,否则这些 episode 不会出现在 `rollout/*` 里。 +- `learn() -> (global_step, log_dict)`:前面调 `pre_learn()`,后面调 + `post_learn()`。基类的 `post_learn` 会清空 episode 指标,覆写时应调用 + `super().post_learn()`。 +- 调度与保存:`should_learn`、`should_save`、`should_stop`、`create_checkpoint` + +另外两个由 CLI 在 server 启动前调用。 + +- `init_optimizers()`:算法构建完立刻调用,默认什么也不做。`dppo` 在这里构建优化器。 +- `load_checkpoint(checkpoint)`:带 `--resume` 启动时调用。 + +还有一个类属性。 + +- `on_policy: bool = True`。只有动作出自当前策略、并且 `should_learn()` 为假的帧, + server 才会交给 `feedback`,其余都交给 `discard_feedback`。从 replay buffer + 学习的算法把 `on_policy` 设为 `False`,就能拿到每一个带 step state 的帧,SAC + 示例就是这样。见[训练循环](ppo.zh.md#what-happens-on-the-server)。 + +这个循环对各个钩子意味着什么。 + +- `should_learn()` 为真时 server 不做推理。如果 `learn()` 之后它仍然为真,server + 会接着再 learn 一次,中间不推理。所以 `learn()` 或 `post_learn()` 必须清空 + buffer,或者推进 `should_learn` 所依据的计数。 +- `should_save()` 和 `should_stop()` 在循环每转一圈时都会检查,而不是每次 learn + 才查一次。`create_checkpoint()` 执行之后 `should_save()` 必须变回假。内置算法都在 + `create_checkpoint` 里记下已保存的轮数。 + +准确签名见 `plugrl_server/algorithm/base_algorithm.py`。server 调用的是 `learn()`,但 `learn()` 在 `BaseAlgorithm` 上已经实现了:它调用 `learn_impl()` -再用 `build_train_info` 包一层结果。抽象方法是 `learn_impl`,要覆写的是它。 -覆写 `learn` 会让 `learn_impl` 悬空、类仍然是抽象的,`make_algo` 会直接抛 -`TypeError`。下面这份模板此前正是这么写的。 +再用 `build_train_info` 包一层结果,`rollout/*` 指标就是在这里加上的。抽象方法是 +`learn_impl`,要覆写的是它。覆写 `learn` 会让 `learn_impl` 悬空、类仍然是抽象的, +`make_algo` 会直接抛 `TypeError`。 `infer` 与 `feedback` 收发的是 `plugrl_server.policy.state` 里的 `PolicyRuntimeState`,`feedback` 还要接 `train_state: PolicyTrainState = None`。 -`plugrl-server` 里没有 `InternalState` 这个类型,也没有 -`get_action_and_internal_state` 这个方法 - 模板里这两个名字都用了,而且都 import -不进来。`feedback` 的每个参数都是 keyword-only,所以名字对不上会在 server 第一次 -调用时直接 `TypeError`,不会被悄悄当成改名放过。 +旧模板里用过 `InternalState` 和 `get_action_and_internal_state`,这两个都不存在。 +`feedback` 的每个参数都是 keyword-only,所以名字对不上会在 server 第一次调用时 +直接 `TypeError`,不会被悄悄当成改名放过。 ## 最小模板 @@ -57,7 +93,7 @@ import numpy as np from plugrl_server.algorithm.base_algorithm import BaseAlgoConfig, BaseAlgorithm from plugrl_server.algorithm.registration import register_algo, register_algo_config from plugrl_server.common.checkpoint_manager import Checkpoint -from plugrl_server.policy.base_policy import BasePolicy +from plugrl_server.policy.base_torch_policy import BaseTorchPolicy from plugrl_server.policy.state import PolicyRuntimeState, PolicyTrainState UID = "your-algo" @@ -66,12 +102,16 @@ UID = "your-algo" @register_algo_config(UID) @dataclasses.dataclass class YourAlgoConfig(BaseAlgoConfig): - total_timesteps: int = 100_000 + # The run length is BaseAlgoConfig's field, set with --algo.global-steps. + global_steps: int | None = 100_000 @register_algo(UID) class YourAlgorithm(BaseAlgorithm): - def __init__(self, config: YourAlgoConfig, policy: BasePolicy): + # Set to False if the algorithm learns from a replay buffer. + on_policy = True + + def __init__(self, config: YourAlgoConfig, policy: BaseTorchPolicy): super().__init__(config=config, policy=policy) self.global_step = 0 @@ -94,6 +134,11 @@ class YourAlgorithm(BaseAlgorithm): next_truncated: bool, prev_node: tuple, ) -> tuple[tuple, int, dict]: + # Store the frame here. + if (next_terminated or next_truncated) and "episode" in info: + if bool(info["episode"].get("mask", True)): + # Without this, rollout/success, reward and length stay 0. + self.record_episode_metrics(info["episode"]) self.global_step += 1 return prev_node, self.global_step, {} @@ -104,12 +149,13 @@ class YourAlgorithm(BaseAlgorithm): return False def should_stop(self) -> bool: - return self.global_step >= self.config.total_timesteps + return self.global_step >= self.config.global_steps def should_save(self) -> bool: return False def create_checkpoint(self) -> Checkpoint: + # state_dict() exists because a BaseTorchPolicy is a torch.nn.Module. return Checkpoint(step=self.global_step, model=self.policy.state_dict(), optimizer=None, meta={}) def load_checkpoint(self, checkpoint: Checkpoint) -> None: @@ -121,6 +167,7 @@ class YourAlgorithm(BaseAlgorithm): ## 设计规则 - 模型结构与动作生成参数放在 policy config。 +- 运行长度用 `BaseAlgoConfig.global_steps`,不要另起一个字段。进度显示读的就是它。 - `learn()` 用外部数据时,显式维护环境步数与更新步数。 - 训练调度相关计数写进 `Checkpoint.meta`,并在 `load_checkpoint` 恢复。 - 不写分布式钩子就用 `BaseAlgorithm`,不要直接上 `DDPAlgorithm`。 @@ -131,7 +178,7 @@ class YourAlgorithm(BaseAlgorithm): ```bash plugrl-run-server dummy-policy default your-algo default -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` 如果算法在外部包里,先 import 再进入 CLI。 @@ -143,7 +190,11 @@ python -c "import my_pkg.plugrl_algorithms; from plugrl_server.cli import main; ## 常见问题 -- `plugrl-run-server --help` 里找不到 UID:模块没有被 import。 +- `plugrl-run-server --help` 里找不到 UID:模块没有被 import。 +- 打印完配置后报 `KeyError: 'Algorithm your-algo is not registered.'`:只 import 了配置模块,类所在的模块没 import。 +- `rollout/*` 一直是 0:`feedback` 没调用 `record_episode_metrics`。如果只是少了一部分 episode,是覆写的 `discard_feedback` 没调用 `super()`。 +- server 反复 learn、不再推理:`learn()` 之后 `should_learn()` 仍然为真。 +- 循环每转一圈都写一次 checkpoint:`create_checkpoint()` 之后 `should_save()` 还是真。 - `--algo.*` 和 `--policy.*` 出现重复语义参数:保留一侧即可。 - resume 后学习频率或保存周期漂移:从 `Checkpoint.meta` 恢复所有计数。 - `DDPAlgorithm` 运行时报错:改用 `BaseAlgorithm`,或补齐分布式钩子。 diff --git a/docs/algorithm/index.md b/docs/algorithm/index.md index 629309b..00e346e 100644 --- a/docs/algorithm/index.md +++ b/docs/algorithm/index.md @@ -13,17 +13,19 @@ plugrl-run-server dppo-policy default dppo hopper ## Verify -List registered algorithms. +List registered algorithms. They are listed one level down, after a policy +and its variant. The top-level `plugrl-run-server --help` lists policy UIDs +only. ```bash -plugrl-run-server --help +plugrl-run-server dummy-policy default --help ``` Run a smoke test loop. ```bash plugrl-run-server dummy-policy default dummy default -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` ## How selection works @@ -40,26 +42,64 @@ Config sources. Discovery. - Registries live in `plugrl_server.policy.registration` and `plugrl_server.algorithm.registration`. -- Your modules must be imported before the CLI is built. +- Your modules must be imported before the CLI is built: the module that + registers the config and the module that registers the class. See + [Custom algorithm](custom_algorithm.md#file-layout). ## Built-in algorithms -Five UIDs are registered in `plugrl-server`. +Six UIDs are registered in `plugrl-server`. -- `fpo`: FPO training loop - the algorithm the quickstarts run -- `dummy`: protocol and connectivity smoke tests +- `fpo`: FPO training loop - the algorithm the quickstarts run. Needs a flow + policy: `fpo-policy` or `pi0-policy`. +- `dppo`: DPPO training loop. Takes a diffusion or a flow policy: + `dppo-policy`, `fpo-policy` or `pi0-policy`. Variants `hopper`, `walker`, + `cheetah`, `square` and `libero` (the one for `pi0-policy`). +- `ppo`: PPO for Gaussian policies, see below. - `eval`: run a policy without training it, optionally from `--algo.policy-checkpoint-path` -- `dppo`: DPPO training loop, needs the `dppo` extra -- `dppo-dist`: distributed DPPO via the Ray launcher, needs the `dppo` extra +- `dummy`: protocol and connectivity smoke tests +- `dppo-dist`: DPPO for the Ray launcher `plugrl-run-server-ray` only. See the + note in [Training loop](ppo.md#ray-launcher). + +No algorithm needs the `dppo` extra. Two policies do: `dppo-policy` and +`dppo-gaussian-policy`. Without the extra, those two policy UIDs are missing +from the CLI, with no error and no warning. Install it with +`uv sync --extra dppo` in `plugrl-server`. + +For `ppo` and `dppo`, a run is `--algo.train-itrs` iterations of +`--algo.buffer-size` frames. Both set `global_steps` from those two, so +`--algo.global-steps` has no effect on them. + +### `ppo` + +PPO as CleanRL's `ppo_continuous_action.py` runs it: clipped surrogate, +clipped value loss, advantages normalized per minibatch, rewards scaled by a +running estimate of the return's deviation. It needs a policy with +`evaluate_actions`, which today means `gaussian-policy` or +`dppo-gaussian-policy`. It refuses a policy run with `--policy.deterministic`. -Without the `dppo` extra installed, `plugrl-run-server` logs -`Could not import DPPO algorithm module` at startup and the `dppo-dist` -subcommand is absent. +Variants. + +- `default`: CleanRL's MuJoCo settings, 488 iterations of 2048 frames. +- `dppo-square`: DPPO's Gaussian PPO baseline on robomimic square, for + `dppo-gaussian-policy` started from DPPO's released checkpoint. + +```bash +# MuJoCo. The default sizes, 17 and 6, fit HalfCheetah and Walker2d; Hopper is 11 and 3. +plugrl-run-server gaussian-policy default ppo default \ + --policy.obs-dim 17 --policy.action-dim 6 + +# robomimic square, from DPPO's released Gaussian checkpoint +plugrl-run-server dppo-gaussian-policy default ppo dppo-square \ + --policy.checkpoint-path /path/to/square_gaussian_pretrained.pt +``` ## Troubleshooting -- Algorithm UID not listed: registration module was not imported. +- Algorithm UID not listed under `plugrl-run-server --help`: registration module was not imported. +- UID listed, but the run dies with `KeyError: 'Algorithm is not registered.'`: the config module was imported and the class module was not. +- `dppo-policy` or `dppo-gaussian-policy` missing from the CLI: the `dppo` extra is not installed. - CLI flags conflict across policy and algo: keep shared concepts in one config. ## Next steps diff --git a/docs/algorithm/index.zh.md b/docs/algorithm/index.zh.md index acb70e5..c07ee24 100644 --- a/docs/algorithm/index.zh.md +++ b/docs/algorithm/index.zh.md @@ -13,17 +13,18 @@ plugrl-run-server dppo-policy default dppo hopper ## 验证 -查看已注册算法。 +查看已注册算法。算法要在选定策略和变体之后的下一层才列出来;顶层的 +`plugrl-run-server --help` 只列策略 UID。 ```bash -plugrl-run-server --help +plugrl-run-server dummy-policy default --help ``` 跑一次 smoke test。 ```bash plugrl-run-server dummy-policy default dummy default -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` ## 选择方式 @@ -40,24 +41,58 @@ plugrl-run-env-client dummy-v1 --num-episodes 1 发现机制。 - registry 位于 `plugrl_server.policy.registration` 与 `plugrl_server.algorithm.registration` -- 模块必须在构建 CLI 前被 import +- 模块必须在构建 CLI 前被 import:注册配置的模块和注册类的模块都要 import, + 见[自定义算法](custom_algorithm.zh.md#file-layout) ## 内置算法 -`plugrl-server` 里一共注册了五个 UID。 +`plugrl-server` 里一共注册了六个 UID。 -- `fpo`:FPO 训练循环 - 快速开始跑的就是它 -- `dummy`:协议与联通性验证 +- `fpo`:FPO 训练循环 - 快速开始跑的就是它。需要流策略:`fpo-policy` 或 `pi0-policy` +- `dppo`:DPPO 训练循环。diffusion 策略和流策略都能用:`dppo-policy`、 + `fpo-policy`、`pi0-policy`。变体有 `hopper`、`walker`、`cheetah`、`square` + 和 `libero`(给 `pi0-policy` 用的) +- `ppo`:面向高斯策略的 PPO,见下文 - `eval`:只跑策略不训练,可用 `--algo.policy-checkpoint-path` 指定权重 -- `dppo`:DPPO 训练循环,需要 `dppo` 可选依赖 -- `dppo-dist`:经 Ray 启动的分布式 DPPO,需要 `dppo` 可选依赖 +- `dummy`:协议与联通性验证 +- `dppo-dist`:只给 Ray 启动器 `plugrl-run-server-ray` 用的 DPPO,见 + [训练循环](ppo.zh.md#ray-launcher)里的说明 + +没有哪个算法需要 `dppo` 可选依赖。需要它的是两个策略:`dppo-policy` 和 +`dppo-gaussian-policy`。没装的话,这两个策略 UID 会直接从 CLI 里消失,既不报错 +也不警告。安装方法是在 `plugrl-server` 里执行 `uv sync --extra dppo`。 + +`ppo` 和 `dppo` 一次运行跑 `--algo.train-itrs` 轮,每轮 `--algo.buffer-size` +帧。两者都用这两个值算出 `global_steps`,所以 `--algo.global-steps` 对它们不起作用。 + +### `ppo` + +按 CleanRL `ppo_continuous_action.py` 的做法实现的 PPO:裁剪的替代目标、裁剪的 +value loss、按 minibatch 归一化的 advantage,reward 按回报标准差的滑动估计缩放。 +它要求策略提供 `evaluate_actions`,目前就是 `gaussian-policy` 和 +`dppo-gaussian-policy`。策略开了 `--policy.deterministic` 时它会拒绝运行。 -没装 `dppo` 可选依赖时,`plugrl-run-server` 启动会打印 -`Could not import DPPO algorithm module`,并且没有 `dppo-dist` 子命令。 +变体。 + +- `default`:CleanRL 的 MuJoCo 设置,488 轮,每轮 2048 帧。 +- `dppo-square`:DPPO 在 robomimic square 上的高斯 PPO 基线,配合从 DPPO 发布的 + checkpoint 起步的 `dppo-gaussian-policy`。 + +```bash +# MuJoCo. The default sizes, 17 and 6, fit HalfCheetah and Walker2d; Hopper is 11 and 3. +plugrl-run-server gaussian-policy default ppo default \ + --policy.obs-dim 17 --policy.action-dim 6 + +# robomimic square, from DPPO's released Gaussian checkpoint +plugrl-run-server dppo-gaussian-policy default ppo dppo-square \ + --policy.checkpoint-path /path/to/square_gaussian_pretrained.pt +``` ## 常见问题 -- `--help` 里找不到 UID:注册模块没有被 import。 +- `plugrl-run-server --help` 里找不到算法 UID:注册模块没有被 import。 +- UID 列出来了,但运行时报 `KeyError: 'Algorithm is not registered.'`:只 import 了配置模块,没 import 类所在的模块。 +- CLI 里没有 `dppo-policy` 或 `dppo-gaussian-policy`:没装 `dppo` 可选依赖。 - policy 与 algo 的 flag 冲突:共享概念只保留在一侧配置。 ## 下一步 diff --git a/docs/algorithm/ppo.md b/docs/algorithm/ppo.md index 0aa1cbb..843ee7a 100644 --- a/docs/algorithm/ppo.md +++ b/docs/algorithm/ppo.md @@ -7,42 +7,124 @@ This page describes what the server loop does at runtime. Run a DPPO experiment. ```bash -plugrl-run-server dppo-policy default dppo hopper --exp_name my_dppo_exp +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp ``` ## Verify -Start one worker and confirm the loop runs. +Run the loop end to end with the dummy pair, which needs no model and no +simulator. ```bash -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-server dummy-policy default dummy default +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` ## What happens on the server The server runs a scheduler loop. -1. Wait until all connected workers have enqueued an `infer` request. -2. Aggregate observations and call `algorithm.infer(batch_obs)`. +1. If `algorithm.should_learn()` is true, do not infer. A full buffer learns first. +2. Otherwise wait until every connected worker has enqueued an `infer` + request, aggregate the observations and call `algorithm.infer(batch_obs)`. + Each action is tagged with the number of learn steps so far. 3. Send `action` back to each worker. -4. Receive `feedback` and call `algorithm.feedback(...)`. -5. Periodically call `algorithm.learn()` and save checkpoints. +4. Receive `feedback` and hand each frame to the algorithm, as described below. +5. Call `algorithm.learn()` when `should_learn()` is true. Save a checkpoint + when `should_save()` or `should_stop()` is true. Shut down after + `should_stop()`. + +Which frames are trained on. + +- A frame goes to `algorithm.feedback(...)` only if its action came from the + current policy (no learn step since it was inferred) and the algorithm is + still collecting (`should_learn()` is false). Only those frames are stored. +- Every other frame goes to + `algorithm.discard_feedback(info=..., next_terminated=..., next_truncated=...)`. + The default records the finished episode's metrics, if the frame ended one, + and does nothing else. +- A frame with no step state is always discarded. That happens when a worker + reconnects mid-run: the server has no previous observation for it. +- For the built-in algorithms, `global_step` counts stored frames only. The + metric `server/discarded_frames` counts the discarded ones since the server + started, and is logged with each learn step. +- This rule is for on-policy algorithms, which is every built-in one + (`BaseAlgorithm.on_policy = True`). An algorithm that learns from a replay + buffer sets `on_policy = False` and is given every frame that has a step + state. `examples/sac/sac.py` does. Key properties. -- Inference is batched by number of active connections. -- Feedback, learning, and saving are guarded by a model lock. +- Inference waits for one request from every active connection, unless + `--mini-infer-batch-size N` is set. Then it runs as soon as N environments + are queued. +- Inference, feedback, learning and saving all run under one model lock, so + they never overlap. +- A worker that sends no feedback for `--feedback-wait-timeout` seconds + (default 60) is disconnected. While a learn step runs, the server keeps + waiting instead. ## Common options - Tracking: `--track.enabled`, `--track.tracker swanlab|wandb` -- Checkpoints: `--checkpoint-base-dir ./checkpoints`, `--resume` +- Checkpoints: `--checkpoint-base-dir ./checkpoints`, `--exp-name NAME`, + `--resume`, `--overwrite` +- Seeding: `--seed 0` seeds Python, NumPy and torch before the policy is built. + A run with one worker is reproducible. With several it is not, because + which requests share a batch depends on arrival order. +- Inference batching: `--mini-infer-batch-size N` +- Worker timeout: `--feedback-wait-timeout 60` + +`--track.enabled`, `--resume` and `--overwrite` are bare boolean flags. +Writing `--track.enabled true` or `--resume true` is a parse error - tyro +reports `Unrecognized arguments: true` and exits. The off switches are +`--track.no-enabled`, `--no-resume` and `--no-overwrite`. + +### Checkpoints and resuming + +- A run writes to `/////`. + Without `--exp-name`, the name is built from a timestamp, so every run gets + a new directory. +- `--resume` loads the latest step in that directory. To reach the same + directory, pass the same `--checkpoint-base-dir` and `--exp-name` and the + same policy and algorithm UIDs. Keep the variants the same too, or the + weights will not fit. If the directory holds no checkpoint, the server exits + with `FileNotFoundError`, and leaves the directory behind, empty. +- If the directory already exists and you pass neither `--resume` nor + `--overwrite`, the server exits with `FileExistsError`. `--overwrite` + deletes the directory first. +- Stopping the server with Ctrl-C writes a checkpoint on the way out, so an + interrupted run can be resumed. It waits up to 30 seconds for a running + learn step, and it does not save after a fatal error. -`--track.enabled` and `--resume` are bare boolean flags. Writing -`--track.enabled true` or `--resume true` is a parse error - tyro reports -`Unrecognized arguments: true` and exits. The off switches are -`--track.no-enabled` and `--no-resume`. Both were written with a `true` -argument on this page. +```bash +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp +# stopped with Ctrl-C; later: +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp --resume +``` + +### Starting from another run's weights + +`fpo` and `dppo` can start from a checkpoint written by another run. +`--algo.policy-checkpoint-path` is a step directory, the one that holds +`model.safetensors`. `--algo.restore` says how much of it to take. + +- `all` (default): model, optimizer, step and iteration. A resume from any + directory. For `dppo` the checkpoint must have been written by `dppo`. +- `model`: weights only. Optimizer, step and iteration start over. +- `except-critic`: every weight except `critic.*`. The value head keeps its + random initialization. + +```bash +plugrl-run-server fpo-policy default fpo default \ + --algo.policy-checkpoint-path ./checkpoints/fpo/fpo-policy/my_fpo_exp/983040 \ + --algo.restore except-critic +``` + +`--algo.restore` other than `all` requires `--algo.policy-checkpoint-path`. +`eval` takes `--algo.policy-checkpoint-path` too, and loads only the weights. + +## Ray launcher {#ray-launcher} Multi GPU runs via the Ray launcher. @@ -54,17 +136,20 @@ plugrl-run-server-ray dppo-policy default dppo-dist hopper --num-ddp-gpus 4 The algorithm has to be `dppo-dist`, not `dppo`: `cli_ray.py` asserts `isinstance(algo, DDPAlgorithm)`, and only `DPPOAlgoDistributed` under the - UID `dppo-dist` mixes `DDPAlgorithm` in. This page previously showed - `dppo`, which trips that assertion. The launcher also requires the `dppo` - extra, builds its worker list from the *local* GPU count - so a multi-node + UID `dppo-dist` mixes `DDPAlgorithm` in. The command above needs the `dppo` + extra because of `dppo-policy`; `dppo-dist` itself does not. The launcher + builds its worker list from the *local* GPU count - so a multi-node cluster still only sees the head node - and its server speaks an older - dialect of the protocol than the WebSocket one. Use `plugrl-run-server` - unless you are working on the Ray path itself. + dialect of the protocol than the WebSocket one. It also predates the frame + rule above: it calls `feedback` on every frame and never discards one. + Use `plugrl-run-server` unless you are working on the Ray path itself. ## Troubleshooting -- `--resume` does nothing: confirm the experiment directory already contains checkpoints. -- Workers hang before the first step: check that every worker reaches the `infer` stage. +- `--resume` fails with `FileNotFoundError`: the directory has no checkpoint. Check `--checkpoint-base-dir`, `--exp-name` and both UIDs. +- `FileExistsError` at startup: the experiment directory exists. Add `--resume` or `--overwrite`, or pick another `--exp-name`. After a failed `--resume` the directory is the empty one that run left. +- Workers hang before the first step: check that every worker reaches the `infer` stage. The server logs `Infer queue has been waiting` every 5 seconds while it waits. +- `server/discarded_frames` rises with each learn step: expected with several workers or environments. The round in which a buffer fills leaves actions in flight, and their frames are discarded. ## Next steps diff --git a/docs/algorithm/ppo.zh.md b/docs/algorithm/ppo.zh.md index 477f33b..9aa7f50 100644 --- a/docs/algorithm/ppo.zh.md +++ b/docs/algorithm/ppo.zh.md @@ -7,41 +7,109 @@ 跑一个 DPPO 实验。 ```bash -plugrl-run-server dppo-policy default dppo hopper --exp_name my_dppo_exp +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp ``` ## 验证 -启动一个 worker,确认闭环在跑。 +用 dummy 这一对把闭环完整跑一遍,不需要模型,也不需要仿真器。 ```bash -plugrl-run-env-client dummy-v1 --num-episodes 1 +plugrl-run-server dummy-policy default dummy default +plugrl-run-env-client dummy-v1 --server-host 127.0.0.1 --server-port 8000 --num-episodes 1 ``` -## server 内部发生了什么 +## server 内部发生了什么 {#what-happens-on-the-server} server 以调度循环驱动训练。 -1. 等待所有已连接 worker 都提交一次 `infer`。 -2. 聚合观测并调用 `algorithm.infer(batch_obs)`。 +1. 如果 `algorithm.should_learn()` 为真,就不做推理:buffer 满了先学习。 +2. 否则等所有已连接 worker 都提交一次 `infer`,聚合观测并调用 + `algorithm.infer(batch_obs)`。每个动作都会记下截至此时做过几次 learn。 3. 向每个 worker 回发 `action`。 -4. 接收 `feedback` 并调用 `algorithm.feedback(...)`。 -5. 在合适时机调用 `algorithm.learn()` 并保存 checkpoint。 +4. 接收 `feedback`,把每一帧按下面的规则交给算法。 +5. `should_learn()` 为真时调用 `algorithm.learn()`;`should_save()` 或 + `should_stop()` 为真时保存 checkpoint;`should_stop()` 之后关闭 server。 + +哪些帧会拿来训练。 + +- 只有同时满足两个条件的帧才会交给 `algorithm.feedback(...)`:它的动作出自当前 + 策略(推理之后没有再 learn 过),并且算法还在收集数据(`should_learn()` 为假)。 + 只有这些帧会被存下来。 +- 其余的帧都交给 + `algorithm.discard_feedback(info=..., next_terminated=..., next_truncated=...)`。 + 默认实现只做一件事:如果这一帧结束了一个 episode,就记下这个 episode 的指标。 +- 没有 step state 的帧一律丢弃。worker 在运行中途重连时会出现这种帧:server + 手上没有它的上一条观测。 +- 对内置算法来说,`global_step` 只数存下来的帧。指标 `server/discarded_frames` + 数的是 server 启动以来丢掉的帧,每次 learn 时随其他指标一起记录。 +- 这条规则针对 on-policy 算法,内置算法全都是(`BaseAlgorithm.on_policy = True`)。 + 从 replay buffer 学习的算法把 `on_policy` 设为 `False`,就会拿到每一个带 + step state 的帧。`examples/sac/sac.py` 就是这么做的。 关键点。 -- 推理批大小由当前连接数决定。 -- feedback、learn、save 在模型锁保护下执行。 +- 推理默认等每个活跃连接各来一个请求;设了 `--mini-infer-batch-size N` 的话, + 排队的环境数凑够 N 个就推理。 +- 推理、feedback、learn、save 都在同一把模型锁下执行,互不重叠。 +- worker 超过 `--feedback-wait-timeout` 秒(默认 60)没发 feedback,连接会被 + 关掉。learn 进行期间 server 会一直等,不按超时处理。 ## 常用参数 - 指标追踪:`--track.enabled`、`--track.tracker swanlab|wandb` -- checkpoint:`--checkpoint-base-dir ./checkpoints`、`--resume` - -`--track.enabled` 与 `--resume` 是不带值的布尔开关。写成 +- checkpoint:`--checkpoint-base-dir ./checkpoints`、`--exp-name NAME`、 + `--resume`、`--overwrite` +- 随机种子:`--seed 0` 在构建策略之前给 Python、NumPy 和 torch 设种子。只连一个 + worker 时结果可复现;连多个时不行,因为哪些请求凑进同一批取决于到达顺序。 +- 推理批大小:`--mini-infer-batch-size N` +- worker 超时:`--feedback-wait-timeout 60` + +`--track.enabled`、`--resume` 与 `--overwrite` 是不带值的布尔开关。写成 `--track.enabled true` 或 `--resume true` 会直接解析失败 - tyro 报 -`Unrecognized arguments: true` 并退出。关掉它们用 `--track.no-enabled` -与 `--no-resume`。本页此前这两处都多写了一个 `true`。 +`Unrecognized arguments: true` 并退出。关掉它们用 `--track.no-enabled`、 +`--no-resume` 与 `--no-overwrite`。 + +### checkpoint 与续训 + +- 一次运行写到 `/////`。 + 不给 `--exp-name` 时名字由时间戳生成,所以每次运行都是新目录。 +- `--resume` 读取该目录下最新的 step。要落到同一个目录,就得给相同的 + `--checkpoint-base-dir`、`--exp-name`,以及相同的策略 UID 和算法 UID。变体也要 + 保持一致,否则权重对不上。目录里没有 checkpoint 时 server 以 + `FileNotFoundError` 退出,并留下一个空目录。 +- 目录已存在、又既没给 `--resume` 也没给 `--overwrite` 时,server 以 + `FileExistsError` 退出。`--overwrite` 会先删掉这个目录。 +- 用 Ctrl-C 停掉 server 时会顺手写一个 checkpoint,中断的运行因此可以续上。它最多 + 等 30 秒让正在进行的 learn 停下;出了致命错误则不保存。 + +```bash +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp +# stopped with Ctrl-C; later: +plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp --resume +``` + +### 从别的运行的权重起步 + +`fpo` 和 `dppo` 可以从另一次运行写的 checkpoint 起步。 +`--algo.policy-checkpoint-path` 指向某个 step 目录,也就是放 +`model.safetensors` 的那个。`--algo.restore` 决定取多少。 + +- `all`(默认):模型、优化器、step 和轮数,相当于从任意目录续训。`dppo` 要求 + checkpoint 也是 `dppo` 写的。 +- `model`:只取权重,优化器、step 和轮数从头开始。 +- `except-critic`:除 `critic.*` 以外的全部权重,value head 保持随机初始化。 + +```bash +plugrl-run-server fpo-policy default fpo default \ + --algo.policy-checkpoint-path ./checkpoints/fpo/fpo-policy/my_fpo_exp/983040 \ + --algo.restore except-critic +``` + +`--algo.restore` 取 `all` 以外的值时必须同时给 `--algo.policy-checkpoint-path`。 +`eval` 也接受 `--algo.policy-checkpoint-path`,只加载权重。 + +## Ray 启动器 {#ray-launcher} 多 GPU 训练使用 Ray 启动。 @@ -53,16 +121,18 @@ plugrl-run-server-ray dppo-policy default dppo-dist hopper --num-ddp-gpus 4 算法必须写 `dppo-dist` 而不是 `dppo`:`cli_ray.py` 里有 `isinstance(algo, DDPAlgorithm)` 断言,而只有 UID 为 `dppo-dist` 的 - `DPPOAlgoDistributed` 混入了 `DDPAlgorithm`。本页此前写的是 `dppo`, - 会卡在这条断言上。这个启动器还需要 `dppo` 可选依赖,它按*本机* GPU 数量 - 构造 worker 列表 - 所以多节点集群也只看得见头节点 - 而且它的 server 说的 - 是比 WebSocket 那套更旧的协议方言。除非你就是在改 Ray 这条路径,否则请用 - `plugrl-run-server`。 + `DPPOAlgoDistributed` 混入了 `DDPAlgorithm`。上面这条命令需要 `dppo` 可选依赖, + 是因为用了 `dppo-policy`;`dppo-dist` 本身不需要。这个启动器按*本机* GPU 数量 + 构造 worker 列表 - 所以多节点集群也只看得见头节点 - 而且它的 server 说的是比 + WebSocket 那套更旧的协议方言。它也早于上面的帧规则:每一帧都调用 `feedback`, + 从不丢帧。除非你就是在改 Ray 这条路径,否则请用 `plugrl-run-server`。 ## 常见问题 -- `--resume` 没生效:确认实验目录里已经有 checkpoint。 -- worker 卡住不 step:检查 worker 是否都走到了 `infer` 阶段。 +- `--resume` 报 `FileNotFoundError`:目录里没有 checkpoint。检查 `--checkpoint-base-dir`、`--exp-name` 和两个 UID。 +- 启动时报 `FileExistsError`:实验目录已存在。加 `--resume` 或 `--overwrite`,或者换一个 `--exp-name`。如果之前 `--resume` 失败过,这个目录就是那次留下的空目录。 +- worker 卡住不 step:检查 worker 是否都走到了 `infer` 阶段。server 等待期间每 5 秒打印一次 `Infer queue has been waiting`。 +- `server/discarded_frames` 每次 learn 都在涨:连着多个 worker 或多个环境时这是正常的。buffer 填满的那一轮总有动作还在路上,它们的帧会被丢掉。 ## 下一步 diff --git a/docs/contributing/algorithm.md b/docs/contributing/algorithm.md index af8da54..bb8eda3 100644 --- a/docs/contributing/algorithm.md +++ b/docs/contributing/algorithm.md @@ -11,8 +11,14 @@ Algorithms live in `plugrl-server`. ## Checklist - Implement the `infer(...)` and `feedback(...)` contract used by the server loop. +- Decide `on_policy`. Keep the default `True` if each learn step may use only + frames the current policy collected. Set `False` if the algorithm learns + from a replay buffer. +- If you override `discard_feedback`, call `super()` so finished episodes are + still recorded. - Implement training and checkpoint hooks as needed. -- Register your config (UID + variants) and make sure it is imported. +- Register your config (UID + variants) and your class, and make sure both + modules are imported. ## Verify @@ -30,7 +36,8 @@ plugrl-run-env-client dummy-v1 --num-episodes 1 --server-host 127.0.0.1 --server ## Troubleshooting -- Algo UID not listed: registration module was not imported. +- Algo UID not listed under `plugrl-run-server --help`: registration module was not imported. +- `KeyError: 'Algorithm is not registered.'`: the config module was imported, the class module was not. - Server crashes on first `infer`: observation schema mismatch. ## Next steps diff --git a/docs/contributing/algorithm.zh.md b/docs/contributing/algorithm.zh.md index 217012c..8306fc4 100644 --- a/docs/contributing/algorithm.zh.md +++ b/docs/contributing/algorithm.zh.md @@ -11,8 +11,11 @@ ## 清单 - 实现 server loop 需要的 `infer(...)` 与 `feedback(...)`。 +- 定好 `on_policy`。每次 learn 只能用当前策略收集的帧,就保留默认的 `True`;从 + replay buffer 学习,就设为 `False`。 +- 如果覆写 `discard_feedback`,要调用 `super()`,已结束的 episode 才会照常记录。 - 按需实现训练与 checkpoint 钩子。 -- 注册配置(UID + variants),并确保模块会被 import。 +- 注册配置(UID + variants)和算法类,并确保这两个模块都会被 import。 ## 验证 @@ -30,7 +33,8 @@ plugrl-run-env-client dummy-v1 --num-episodes 1 --server-host 127.0.0.1 --server ## 常见问题 -- CLI 找不到 UID:注册模块没有被 import。 +- `plugrl-run-server --help` 里找不到算法 UID:注册模块没有被 import。 +- 报 `KeyError: 'Algorithm is not registered.'`:import 了配置模块,没 import 类所在的模块。 - 首次 `infer` 崩溃:观测 schema 对不上。 ## 下一步 diff --git a/docs/contributing/policy.md b/docs/contributing/policy.md index d44c624..df78a7e 100644 --- a/docs/contributing/policy.md +++ b/docs/contributing/policy.md @@ -12,6 +12,9 @@ Policies live in `plugrl-server`. - Implement a policy and its config dataclass. - Register it (UID + variants) and make sure it is imported. - Keep action shapes stable across inference and training. +- Return actions batch-first with a horizon axis, `(B, H, ...)`, and set + `action_dim` and `action_horizon` in `__init__`. See the + [contract](../policy/custom_policy.md#contract). ## Verify @@ -30,7 +33,7 @@ plugrl-run-env-client dummy-v1 --num-episodes 1 --server-host 127.0.0.1 --server ## Troubleshooting - Policy UID not listed: registration module was not imported. -- Shape mismatch in training: align action and `InternalState` with your buffers. +- Shape mismatch in training: align the action, the runtime state (`PolicyRuntimeState`) and the train state your algorithm derives from it with your buffers. ## Next steps diff --git a/docs/contributing/policy.zh.md b/docs/contributing/policy.zh.md index 4bfb9cb..5ef4af4 100644 --- a/docs/contributing/policy.zh.md +++ b/docs/contributing/policy.zh.md @@ -12,6 +12,8 @@ - 实现策略与配置 dataclass。 - 注册(UID + variants),并确保模块会被 import。 - 动作形状在推理与训练路径保持一致。 +- 动作按 batch 在前、带 horizon 轴的 `(B, H, ...)` 返回,并在 `__init__` 里设置 + `action_dim` 和 `action_horizon`。见[约定](../policy/custom_policy.zh.md#contract)。 ## 验证 @@ -30,7 +32,7 @@ plugrl-run-env-client dummy-v1 --num-episodes 1 --server-host 127.0.0.1 --server ## 常见问题 - CLI 找不到 UID:注册模块没有被 import。 -- 训练 shape 对不上:动作与 `InternalState` 结构要和 buffer 对齐。 +- 训练 shape 对不上:动作、runtime state(`PolicyRuntimeState`)以及算法由它导出的 train state,结构都要和 buffer 对齐。 ## 下一步 diff --git a/docs/policy/custom_policy.md b/docs/policy/custom_policy.md index d2ffec9..c40fe75 100644 --- a/docs/policy/custom_policy.md +++ b/docs/policy/custom_policy.md @@ -29,6 +29,11 @@ class YourPolicyConfig(BasePolicyConfig): @register_policy(UID) class YourPolicy(BasePolicy): + def __init__(self, config: YourPolicyConfig): + super().__init__(config) + self.action_dim = ... + self.action_horizon = ... + def prepare_observation(self, obs: dict): ... @@ -39,6 +44,11 @@ class YourPolicy(BasePolicy): ... ``` +For a torch model, subclass `BaseTorchPolicy` and `BaseTorchPolicyConfig` +from `plugrl_server.policy.base_torch_policy` instead. The policy is then a +`torch.nn.Module`, so it has the `state_dict()` that checkpoints need, and it +gets a `--policy.device` flag. + Reference implementation: `plugrl-server/examples/sac/sac_policy.py`. ## File layout @@ -74,7 +84,19 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma an `np.ndarray` or a nested mapping of them. `BaseTorchPolicy` converts that to tensors for you in `extract_model_obs_tensor`. - `get_action_and_runtime_state` returns an action and a `PolicyRuntimeState`. +- The action is batch-first with a horizon axis, `(B, H, ...)`. The server + swaps the first two axes before it sends actions to a client. A policy that + acts one step at a time returns `(B, 1, D)`, as `gaussian-policy` does with + `action[:, None, :]`. +- Set `action_dim` and `action_horizon` in `__init__`. The server sends both + to clients in its metadata message. A policy without them still runs, but + the message leaves them out and every client has to be given the shape by + hand. +- The runtime state must be indexable by batch. The server slices it per + environment: a dataclass or a dict field by field, anything else as + `state[i:i+1]`. Arrays, tensors and `None` all work. - `fake_runtime_state` returns shapes and dtypes that match your buffers. + Algorithms call it through `example_train_state` to size them. - `PolicyRuntimeState` is a type alias, not a base class: return whatever your algorithm needs - a dict, a dataclass, or `None`. diff --git a/docs/policy/custom_policy.zh.md b/docs/policy/custom_policy.zh.md index 7454589..de2ef14 100644 --- a/docs/policy/custom_policy.zh.md +++ b/docs/policy/custom_policy.zh.md @@ -29,6 +29,11 @@ class YourPolicyConfig(BasePolicyConfig): @register_policy(UID) class YourPolicy(BasePolicy): + def __init__(self, config: YourPolicyConfig): + super().__init__(config) + self.action_dim = ... + self.action_horizon = ... + def prepare_observation(self, obs: dict): ... @@ -39,6 +44,10 @@ class YourPolicy(BasePolicy): ... ``` +如果是 torch 模型,改为继承 `plugrl_server.policy.base_torch_policy` 里的 +`BaseTorchPolicy` 和 `BaseTorchPolicyConfig`。这样策略本身就是 `torch.nn.Module`, +有保存 checkpoint 需要的 `state_dict()`,还会多一个 `--policy.device` 参数。 + 参考实现:`plugrl-server/examples/sac/sac_policy.py`。 ## 代码放哪里 @@ -68,12 +77,21 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma your-policy default dummy default ``` -## 约定 +## 约定 {#contract} - `prepare_observation` 把 worker 观测 dict 转成 `NumpyState`,也就是 `np.ndarray` 或它们的嵌套 mapping;转张量由 `BaseTorchPolicy.extract_model_obs_tensor` 负责 - `get_action_and_runtime_state` 返回动作与 `PolicyRuntimeState` -- `fake_runtime_state` 返回能用于 buffer 预分配的形状与 dtype +- 动作是 batch 在前、带 horizon 轴的 `(B, H, ...)`。server 发给 client 之前会交换前 + 两个轴。一次只出一步动作的策略返回 `(B, 1, D)`,`gaussian-policy` 就是用 + `action[:, None, :]` 做到的 +- 在 `__init__` 里设置 `action_dim` 和 `action_horizon`。server 会把这两个值放进 + metadata 消息发给 client。不设也能跑,但消息里就没有这两项,每个 client 都得 + 手工告诉它动作形状 +- runtime state 必须能按 batch 切片。server 会按环境把它切开:dataclass 和 dict + 逐字段切,其他类型按 `state[i:i+1]` 切。数组、张量和 `None` 都可以 +- `fake_runtime_state` 返回能用于 buffer 预分配的形状与 dtype,算法通过 + `example_train_state` 调用它来确定 buffer 大小 - `PolicyRuntimeState` 是类型别名而不是基类:dict、dataclass 或 `None` 都可以, 按算法需要返回 diff --git a/docs/policy/dppo_policy.md b/docs/policy/dppo_policy.md index bd4c1ff..1cb7710 100644 --- a/docs/policy/dppo_policy.md +++ b/docs/policy/dppo_policy.md @@ -1,12 +1,14 @@ # DPPO policies -Use diffusion-style server policies with the DPPO algorithm. +Use diffusion-style server policies with the DPPO algorithm, and flow +policies with DPPO or FPO. This page covers: - Built-in `dppo-policy` - Built-in `pi0-policy` (OpenPI PI0) -- Implementing a diffusion-style policy with `BasePolicyGradientDiffusionPolicy` +- Implementing a diffusion-style policy with `BasePolicyGradientDiffusionPolicy`, + or a flow policy with `BasePolicyGradientFlowPolicy` ## Quickstart @@ -16,9 +18,13 @@ DPPO policy. plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp ``` -OpenPI PI0 policy (requires a checkpoint directory). It is a flow policy: it -pairs with `fpo` or `eval`, never with `dppo`. These are the two invocations -[E11](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e11-vla-rl-libero) ran. +OpenPI PI0 policy (requires a checkpoint directory). It is a flow policy and +runs with `eval`, `fpo` and `dppo`. The first two commands are the ones +[E11](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e11-vla-rl-libero) +ran, without its port, logging and output-directory flags. The third has the +algorithm settings of +[E25](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e25-pi0-dppo), +which ran it end to end. ```bash # evaluation - no learning @@ -33,16 +39,26 @@ plugrl-run-server pi0-policy default fpo default \ --policy.checkpoint-path /path/to/pi0_checkpoint \ --policy.device cuda \ --algo.learning-rate 1e-5 --algo.batch-size 8 \ - --algo.n-samples-per-action 4 --algo.buffer-size 4096 \ - --algo.global-steps 40960 + --algo.num-updates-per-batch 4 --algo.n-samples-per-action 4 \ + --algo.buffer-size 4096 --algo.clipping-epsilon 0.05 \ + --algo.global-steps 40960 --algo.save-interval 1 + +# DPPO fine-tuning +plugrl-run-server pi0-policy default dppo libero \ + --policy.name pi05_libero \ + --policy.checkpoint-path /path/to/pi0_checkpoint \ + --policy.device cuda \ + --algo.buffer-size 4096 --algo.batch-size 8 \ + --algo.train-itrs 2 --algo.save-interval 1 ``` ## Verify -List registered policies and variants. +List registered policies, then the variants of one. ```bash plugrl-run-server --help +plugrl-run-server dppo-policy --help ``` Smoke-test a plug-in policy UID. @@ -56,11 +72,21 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma - UID: `dppo-policy` - Code: `plugrl-server/src/plugrl_server/policy/dppo/dppo_policy.py` -- Dependency: `dppo` must be installed in the server environment. +- Dependency: the `dppo` extra must be installed in the server environment. + Without it, `dppo-policy` is missing from the CLI. - uv: in `plugrl-server`, run `uv sync --extra dppo`. +- Variants, each meant for the `dppo` variant of the same name: + - `default` and `hopper`: gym `hopper-medium-v2` + - `walker`: gym `walker2d-medium-v2` + - `cheetah`: gym `halfcheetah-medium-v2` + - `square`: robomimic `square`, fine-tuning the last 10 of its 20 + denoising steps - Loads: - `plugrl_server/meta/dppo/cfg//.yaml` - `plugrl_server/meta/dppo/asset///normalization.npz` +- `--policy.checkpoint-path` loads a DPPO pretraining checkpoint (`.pt`), + its `ema` weights when it has them. Without it the network starts from + random initialization. - Expects worker obs to contain `states` and concatenates `low_dim_keys`. Common flags. @@ -68,22 +94,31 @@ Common flags. - `--policy.env-type gym` - `--policy.env-name hopper-medium-v2` - `--policy.checkpoint-path /path/to/checkpoint.pt` +- `--policy.ft-denoising-steps 10`: fine-tune only the chain's last 10 + denoising steps. The earlier steps run a frozen copy of the loaded network, + and only the last 10 are recorded for training. Unset, every step is + fine-tuned. `square` sets 10. - `--policy.critic.*` ## Built-in: `pi0-policy` (OpenPI) - UID: `pi0-policy` - Code: `plugrl-server/src/plugrl_server/policy/openpi/openpi_policy.py` +- A flow policy (`BasePolicyGradientFlowPolicy`). Runs with `eval`, `fpo`, + and `dppo` through the `libero` variant. - `--policy.checkpoint-path` is required and must point to a directory with: - `model.safetensors` - `assets/` with normalization stats - Setup notes: `plugrl-server/src/plugrl_server/policy/openpi/README.md`. + Without OpenPI installed, `pi0-policy` is missing from the CLI. Common flags. - `--policy.name pi05_libero` - `--policy.denoising-steps 5` -- `--policy.train-expert-only true` +- `--policy.train-expert-only`: train only the action expert, with the + vision-language model frozen. This is the default; + `--policy.no-train-expert-only` trains both. - `--policy.default-prompt "..."` ## Implement a diffusion-style policy @@ -93,11 +128,23 @@ Subclass `BasePolicyGradientDiffusionPolicy` in The base class drives the denoising loop and fills a `DiffusionRuntimeState`: -- Per-step: `action`, `logprob`, `entropy`, plus `obs["x"]` and `obs["t"]`. -- Final: calls `_postprocess_action` and stores `value` into `runtime_state.value`. +- `runtime_state.obs` is a `DiffusionObs` dataclass with fields `x`, `t` and `cond`. +- Per recorded step: `action`, `logprob`, `entropy`, plus `obs.x` and `obs.t`. + Only the last `num_recorded_denoising_steps` steps are recorded. That is + every step, unless the policy overrides the property, as `dppo-policy` does + for `--policy.ft-denoising-steps`. +- Final: calls `_postprocess_action`, stores the model observation in + `obs.cond` and the value in `runtime_state.value`. What you implement. +- In `__init__`, set `action_dim`, `action_horizon` and `num_denoising_steps`. + `fake_runtime_state` builds its shapes from them, and the server sends the + first two to clients in its metadata message. +- An `actor` module and a `critic` module (the value head). `dppo` raises + `DPPO requires a critic.` without one, and it builds its optimizers from + `policy.actor` and `policy.critic`. `fpo` optimizes the same two. +- `prepare_observation`, which turns the worker obs into arrays - `_get_timesteps`, `_initialize_x`, `_denoising_step`, `_iterative_process_action`, `_postprocess_action` - `fake_diffusion_cond` for buffer preallocation - `_get_value`. It is not abstract, but the base assigns its result into @@ -107,10 +154,16 @@ What you implement. once per inference and passes the result to `_denoising_step` as the keyword-only `cond_cache=` and to `_get_value` as `obs_cache=`. -Key shapes (from `fake_runtime_state`). +At training time `dppo` calls `_denoising_step` again, with `x` of batch +`B * S` (one entry per recorded step), `cond` of batch `B`, and no +`cond_cache`. Expand `cond` with `repeat_interleave` when the two differ, as +`dppo-policy` does. + +Key shapes (from `fake_runtime_state`), with `S = num_recorded_denoising_steps`, +`H = action_horizon`, `D = action_dim`. -- `action/logprob/entropy/obs["x"]`: `(B, S, H, D)` -- `obs["t"]`: `(B, S)` +- `action`, `logprob`, `entropy`, `obs.x`: `(B, S, H, D)` +- `obs.t`: `(B, S)` - `value`: `(B,)` Minimal template. @@ -142,7 +195,11 @@ class MyDPPOPolicyConfig(BasePolicyGradientDiffusionPolicyConfig): class MyDPPOPolicy(BasePolicyGradientDiffusionPolicy): def __init__(self, config: MyDPPOPolicyConfig): super().__init__(config) - ... + self.actor = ... # torch.nn.Module + self.critic = ... # torch.nn.Module, the value head + self.action_dim = ... + self.action_horizon = ... + self.num_denoising_steps = ... def prepare_observation(self, _obs: dict) -> dict[str, np.ndarray]: ... @@ -178,6 +235,24 @@ class MyDPPOPolicy(BasePolicyGradientDiffusionPolicy): ... ``` +### Flow policies + +`BasePolicyGradientFlowPolicy`, in `base_policy_gradient_flow_policy.py`, +subclasses the diffusion base and implements `_denoising_step` for you as an +Euler step along a predicted velocity. `fpo` requires a policy built on it. +A subclass: + +- implements `_predict_v(x, t, cond, *, cond_cache=None)`, returning the + velocity, instead of `_denoising_step`; +- sets `dt` in `__init__`. The built-ins run from `t = 1` to `t = 0` with + `dt = -1.0 / num_denoising_steps`; +- implements the rest of the list above. + +Without a sampling noise level the step is deterministic and its +log-probability is zero. `dppo` passes one, which makes the step stochastic, +and that is how `dppo` trains a flow policy. `fpo-policy` and `pi0-policy` are +the examples to read. + ## Example: LeRobot diffusion adapter See `plugrl-server/examples/lerobot/lerobot_diffusion.py` (`UID = "lerobot-diffusion-policy"`). @@ -187,14 +262,15 @@ Patterns to copy. - Pack multi-step observations into one batched nested `dict` of arrays - that is what `TorchTree` is; the base converts it to tensors for you. - Cache encoder outputs in `build_obs_cache`. -- Support expanded batch `B * num_denoising_steps` with `repeat_interleave`. +- Support the expanded batch `B * S` with `repeat_interleave`, as described above. - Deterministic sampling can return zero `logprob` like `Pi0Policy`. - Stochastic sampling should compute `logprob/entropy` like `DPPOPolicy`. ## Troubleshooting -- `dppo-policy` import fails: install `dppo` in the environment that runs `plugrl-run-server`. -- `pi0-policy` import/setup fails: follow the OpenPI setup in the server repo. +- `dppo-policy` missing from the CLI: install the `dppo` extra in the environment that runs `plugrl-run-server`. +- `pi0-policy` missing from the CLI, or failing at startup: follow the OpenPI setup in the server repo. +- `DPPO requires a critic.`: set `self.critic` in `__init__`. - Shape mismatch in training: keep action shapes stable and align the runtime state with your buffers. - Policy UID not listed: your registration module was not imported. diff --git a/docs/policy/dppo_policy.zh.md b/docs/policy/dppo_policy.zh.md index 9149d3d..642c0bf 100644 --- a/docs/policy/dppo_policy.zh.md +++ b/docs/policy/dppo_policy.zh.md @@ -1,12 +1,13 @@ # DPPO 策略 -在 server 侧运行 diffusion 风格策略,并与 DPPO 算法配合。 +在 server 侧运行 diffusion 风格策略并与 DPPO 算法配合;流策略则可以配 DPPO 或 FPO。 本页包含: - 内置 `dppo-policy` - 内置 `pi0-policy`(OpenPI PI0) -- 用 `BasePolicyGradientDiffusionPolicy` 实现自定义 diffusion policy +- 用 `BasePolicyGradientDiffusionPolicy` 实现自定义 diffusion policy,或用 + `BasePolicyGradientFlowPolicy` 实现流策略 ## 快速开始 @@ -16,32 +17,46 @@ DPPO 策略。 plugrl-run-server dppo-policy default dppo hopper --exp-name my_dppo_exp ``` -OpenPI PI0 策略(需要 checkpoint 目录)。它是流策略:与 `fpo` 或 `eval` 搭配, -**绝不**与 `dppo` 搭配。下面两条就是 [E11](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e11-vla-rl-libero) 实际跑的命令。 +OpenPI PI0 策略(需要 checkpoint 目录)。它是流策略,可以配 `eval`、`fpo` 和 +`dppo`。前两条是 +[E11](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e11-vla-rl-libero) +实际跑的命令,去掉了端口、日志和输出目录相关的参数。第三条用的是 +[E25](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e25-pi0-dppo) +的算法设置,E25 用它完整跑通过。 ```bash -# 评测 —— 不训练 +# evaluation - no learning plugrl-run-server pi0-policy default eval default \ --policy.name pi05_libero \ --policy.checkpoint-path /path/to/pi0_checkpoint \ --policy.device cuda -# FPO 微调 +# FPO fine-tuning plugrl-run-server pi0-policy default fpo default \ --policy.name pi05_libero \ --policy.checkpoint-path /path/to/pi0_checkpoint \ --policy.device cuda \ --algo.learning-rate 1e-5 --algo.batch-size 8 \ - --algo.n-samples-per-action 4 --algo.buffer-size 4096 \ - --algo.global-steps 40960 + --algo.num-updates-per-batch 4 --algo.n-samples-per-action 4 \ + --algo.buffer-size 4096 --algo.clipping-epsilon 0.05 \ + --algo.global-steps 40960 --algo.save-interval 1 + +# DPPO fine-tuning +plugrl-run-server pi0-policy default dppo libero \ + --policy.name pi05_libero \ + --policy.checkpoint-path /path/to/pi0_checkpoint \ + --policy.device cuda \ + --algo.buffer-size 4096 --algo.batch-size 8 \ + --algo.train-itrs 2 --algo.save-interval 1 ``` ## 验证 -查看注册到 CLI 的策略与变体。 +先列出已注册的策略,再看某个策略有哪些变体。 ```bash plugrl-run-server --help +plugrl-run-server dppo-policy --help ``` 外部包策略先 import 再进入 CLI。 @@ -55,11 +70,18 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma - UID:`dppo-policy` - 代码:`plugrl-server/src/plugrl_server/policy/dppo/dppo_policy.py` -- 依赖:运行 `plugrl-run-server` 的环境里需要安装 `dppo`。 +- 依赖:server 环境里要装 `dppo` 可选依赖。没装的话 CLI 里就没有 `dppo-policy`。 - uv:在 `plugrl-server` 执行 `uv sync --extra dppo`。 +- 变体,各自配同名的 `dppo` 变体: + - `default` 和 `hopper`:gym `hopper-medium-v2` + - `walker`:gym `walker2d-medium-v2` + - `cheetah`:gym `halfcheetah-medium-v2` + - `square`:robomimic `square`,只微调 20 个去噪步里的最后 10 个 - 加载: - `plugrl_server/meta/dppo/cfg//.yaml` - `plugrl_server/meta/dppo/asset///normalization.npz` +- `--policy.checkpoint-path` 加载 DPPO 预训练 checkpoint(`.pt`),有 `ema` 权重时用 + `ema`。不给的话网络从随机初始化开始。 - 观测:期望 worker 观测包含 `states`,并按 `low_dim_keys` 拼接。 常用参数。 @@ -67,22 +89,29 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma - `--policy.env-type gym` - `--policy.env-name hopper-medium-v2` - `--policy.checkpoint-path /path/to/checkpoint.pt` +- `--policy.ft-denoising-steps 10`:只微调去噪链最后 10 步。前面的步骤用加载进来的 + 网络的冻结副本跑,只有最后 10 步会被记录下来用于训练。不设则每一步都微调。 + `square` 设的就是 10。 - `--policy.critic.*` ## 内置:`pi0-policy`(OpenPI) - UID:`pi0-policy` - 代码:`plugrl-server/src/plugrl_server/policy/openpi/openpi_policy.py` +- 是流策略(`BasePolicyGradientFlowPolicy`),可以配 `eval`、`fpo`,也可以通过 + `libero` 变体配 `dppo`。 - `--policy.checkpoint-path` 必填,目录内需要: - `model.safetensors` - `assets/`(归一化统计) - 本地安装/替换步骤见:`plugrl-server/src/plugrl_server/policy/openpi/README.md`。 + 没装 OpenPI 时 CLI 里没有 `pi0-policy`。 常用参数。 - `--policy.name pi05_libero` - `--policy.denoising-steps 5` -- `--policy.train-expert-only true` +- `--policy.train-expert-only`:只训练 action expert,视觉语言模型冻结。这是默认值; + `--policy.no-train-expert-only` 两者一起训练。 - `--policy.default-prompt "..."` ## 自定义 diffusion policy @@ -92,11 +121,22 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma 基类负责 denoising 循环并填充 `DiffusionRuntimeState`: -- 每步写入:`action`、`logprob`、`entropy`,以及 `obs["x"]`、`obs["t"]`。 -- 最终调用 `_postprocess_action`,并把 `value` 写入 `runtime_state.value`。 +- `runtime_state.obs` 是一个 `DiffusionObs` dataclass,字段为 `x`、`t`、`cond`。 +- 每个被记录的步骤写入:`action`、`logprob`、`entropy`,以及 `obs.x`、`obs.t`。 + 只记录最后 `num_recorded_denoising_steps` 步。默认就是全部步骤,除非策略覆写这个 + 属性,`dppo-policy` 为了 `--policy.ft-denoising-steps` 就覆写了它。 +- 最终调用 `_postprocess_action`,把模型观测存进 `obs.cond`,把 value 写入 + `runtime_state.value`。 你需要实现。 +- 在 `__init__` 里设置 `action_dim`、`action_horizon` 和 `num_denoising_steps`。 + `fake_runtime_state` 按它们构造形状,server 也会把前两个放进 metadata 消息发给 + client。 +- 一个 `actor` 模块和一个 `critic` 模块(value head)。没有 critic 时 `dppo` 会报 + `DPPO requires a critic.`,它的优化器也是从 `policy.actor` 和 `policy.critic` + 建的。`fpo` 优化的也是这两个。 +- `prepare_observation`,把 worker 观测转成数组 - `_get_timesteps`、`_initialize_x`、`_denoising_step`、`_iterative_process_action`、`_postprocess_action` - `fake_diffusion_cond`(用于 buffer 预分配) - `_get_value`。它不是 abstract,但基类会把它的返回值写进 `runtime_state.value`, @@ -106,10 +146,15 @@ python -c "import my_pkg.plugrl_policies; from plugrl_server.cli import main; ma 结果以 keyword-only 的 `cond_cache=` 传给 `_denoising_step`, 以 `obs_cache=` 传给 `_get_value`。 -关键形状(来自 `fake_runtime_state`)。 +训练时 `dppo` 会再调用 `_denoising_step`,这时 `x` 的 batch 是 `B * S`(每个被记录 +的步骤一条),`cond` 的 batch 是 `B`,而且不传 `cond_cache`。两者不一致时用 +`repeat_interleave` 扩展 `cond`,`dppo-policy` 就是这么做的。 + +关键形状(来自 `fake_runtime_state`),其中 `S = num_recorded_denoising_steps`、 +`H = action_horizon`、`D = action_dim`。 -- `action/logprob/entropy/obs["x"]`:`(B, S, H, D)` -- `obs["t"]`:`(B, S)` +- `action`、`logprob`、`entropy`、`obs.x`:`(B, S, H, D)` +- `obs.t`:`(B, S)` - `value`:`(B,)` 最小模板。 @@ -141,7 +186,11 @@ class MyDPPOPolicyConfig(BasePolicyGradientDiffusionPolicyConfig): class MyDPPOPolicy(BasePolicyGradientDiffusionPolicy): def __init__(self, config: MyDPPOPolicyConfig): super().__init__(config) - ... + self.actor = ... # torch.nn.Module + self.critic = ... # torch.nn.Module, the value head + self.action_dim = ... + self.action_horizon = ... + self.num_denoising_steps = ... def prepare_observation(self, _obs: dict) -> dict[str, np.ndarray]: ... @@ -177,6 +226,20 @@ class MyDPPOPolicy(BasePolicyGradientDiffusionPolicy): ... ``` +### 流策略 + +`base_policy_gradient_flow_policy.py` 里的 `BasePolicyGradientFlowPolicy` 继承自 +diffusion 基类,并替你实现了 `_denoising_step`:沿预测的速度场走一个 Euler 步。 +`fpo` 要求策略基于它。子类需要: + +- 实现 `_predict_v(x, t, cond, *, cond_cache=None)`,返回速度,代替 `_denoising_step`; +- 在 `__init__` 里设置 `dt`。内置策略从 `t = 1` 走到 `t = 0`,用 + `dt = -1.0 / num_denoising_steps`; +- 实现上面清单里的其余部分。 + +不给采样噪声时这一步是确定性的,log-probability 为 0。`dppo` 会传入噪声,让这一步 +变成随机的,`dppo` 正是这样训练流策略的。可以参考 `fpo-policy` 和 `pi0-policy`。 + ## LeRobot 示例 参考 `plugrl-server/examples/lerobot/lerobot_diffusion.py`(`UID = "lerobot-diffusion-policy"`)。 @@ -185,14 +248,15 @@ class MyDPPOPolicy(BasePolicyGradientDiffusionPolicy): - 把多步观测打包成一个 batched 的嵌套 `dict`(这就是 `TorchTree`),基类会替你转成张量。 - 在 `build_obs_cache` 缓存 encoder 输出。 -- 用 `repeat_interleave` 支持 `B * num_denoising_steps` 的扩展 batch。 +- 用 `repeat_interleave` 支持上面说的 `B * S` 扩展 batch。 - 确定性采样可像 `Pi0Policy` 一样返回全 0 的 `logprob`。 - 随机采样按分布计算 `logprob/entropy`,与 `DPPOPolicy` 对齐。 ## 常见问题 -- `dppo-policy` import 失败:在运行 `plugrl-run-server` 的环境里安装 `dppo`。 -- `pi0-policy` 启动失败:按 OpenPI README 完成本地设置。 +- CLI 里没有 `dppo-policy`:在运行 `plugrl-run-server` 的环境里安装 `dppo` 可选依赖。 +- CLI 里没有 `pi0-policy`,或它启动失败:按 OpenPI README 完成本地设置。 +- 报 `DPPO requires a critic.`:在 `__init__` 里设置 `self.critic`。 - 训练 shape 对不上:动作形状要稳定,runtime state 字段要与 buffer 对齐。 - CLI 找不到 UID:注册模块没有被 import。 diff --git a/docs/policy/index.md b/docs/policy/index.md index c367b2e..17e0e7e 100644 --- a/docs/policy/index.md +++ b/docs/policy/index.md @@ -19,13 +19,35 @@ plugrl-run-server --help ## Built-in policies -- `dummy-policy`: random actions for protocol smoke tests -- `dppo-policy`: DPPO diffusion policy +- `dummy-policy`: random actions for protocol smoke tests. By default they are + continuous, 7-dimensional, with a horizon of 4, as the env client's + `dummy-v1` expects. - `fpo-policy`: FPO flow-matching policy, used by the get-started run -- `pi0-policy`: OpenPI policy, requires a checkpoint path +- `dppo-policy`: DPPO diffusion policy. Needs the `dppo` extra. +- `gaussian-policy`: CleanRL's Gaussian MLP, trained with `ppo`. + `--policy.deterministic` acts with the mean instead of a sample, for + evaluation with `eval`. +- `dppo-gaussian-policy`: DPPO's Gaussian MLP, which loads the checkpoints + DPPO releases. Trained with `ppo`. Needs the `dppo` extra. +- `pi0-policy`: OpenPI policy, requires a checkpoint path. Needs OpenPI. -OpenPI example. `pi0-policy` is a flow policy, so it pairs with `fpo` or with -`eval` - not with `dppo`, which expects a diffusion policy. +Which algorithm takes which policy. + +- `fpo`: flow policies, `fpo-policy` and `pi0-policy`. +- `dppo`: diffusion policies and flow policies, since the flow base class + subclasses the diffusion one: `dppo-policy`, `fpo-policy`, `pi0-policy`. +- `ppo`: `gaussian-policy`, `dppo-gaussian-policy`. +- `eval` and `dummy`: any policy. + +A policy whose dependencies are missing is also missing from the CLI, with no +error. `dppo-policy` and `dppo-gaussian-policy` need the `dppo` extra +(`uv sync --extra dppo` in `plugrl-server`). `pi0-policy` needs OpenPI, set up +as `plugrl-server/src/plugrl_server/policy/openpi/README.md` describes. + +OpenPI example. `pi0-policy` runs with `eval`, with `fpo`, and with `dppo` +through its `libero` variant, which +[E25](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e25-pi0-dppo) +ran end to end. ```bash plugrl-run-server pi0-policy default eval default \ @@ -36,8 +58,8 @@ plugrl-run-server pi0-policy default eval default \ ## Troubleshooting -- Policy UID not listed: registration module was not imported. -- OpenPI import fails: complete the OpenPI local setup in the server repo. +- Policy UID not listed: registration module was not imported, or the policy's dependencies are not installed (see above). +- `pi0-policy` fails at startup: complete the OpenPI local setup in the server repo. ## Next steps diff --git a/docs/policy/index.zh.md b/docs/policy/index.zh.md index 88239f5..0f537b5 100644 --- a/docs/policy/index.zh.md +++ b/docs/policy/index.zh.md @@ -19,13 +19,32 @@ plugrl-run-server --help ## 内置策略 -- `dummy-policy`:随机动作,用于协议联通性验证 -- `dppo-policy`:DPPO diffusion 策略 +- `dummy-policy`:随机动作,用于协议联通性验证。默认是连续动作、7 维、horizon 为 4, + 正好是 env client 的 `dummy-v1` 要的形状 - `fpo-policy`:FPO flow matching 策略,快速上手那条命令用的就是它 -- `pi0-policy`:OpenPI 策略,需要 checkpoint 路径 +- `dppo-policy`:DPPO diffusion 策略,需要 `dppo` 可选依赖 +- `gaussian-policy`:CleanRL 的高斯 MLP,用 `ppo` 训练。`--policy.deterministic` + 让它直接输出均值而不采样,配合 `eval` 做评测 +- `dppo-gaussian-policy`:DPPO 的高斯 MLP,能加载 DPPO 发布的 checkpoint,用 `ppo` + 训练,需要 `dppo` 可选依赖 +- `pi0-policy`:OpenPI 策略,需要 checkpoint 路径,需要装好 OpenPI -OpenPI 示例。`pi0-policy` 是流策略,因此与 `fpo` 或 `eval` 搭配, -**不能**与 `dppo` 搭配——后者要的是 diffusion 策略。 +各算法能配哪些策略。 + +- `fpo`:流策略,即 `fpo-policy` 和 `pi0-policy`。 +- `dppo`:diffusion 策略和流策略都行,因为流策略的基类是 diffusion 基类的子类: + `dppo-policy`、`fpo-policy`、`pi0-policy`。 +- `ppo`:`gaussian-policy`、`dppo-gaussian-policy`。 +- `eval` 和 `dummy`:任何策略。 + +依赖没装全的策略也不会出现在 CLI 里,而且不报错。`dppo-policy` 和 +`dppo-gaussian-policy` 需要 `dppo` 可选依赖(在 `plugrl-server` 里执行 +`uv sync --extra dppo`)。`pi0-policy` 需要 OpenPI,按 +`plugrl-server/src/plugrl_server/policy/openpi/README.md` 安装。 + +OpenPI 示例。`pi0-policy` 可以配 `eval`、`fpo`,也可以通过 `libero` 变体配 +`dppo`,[E25](https://github.com/PlugRL/plugrl-server/tree/main/experiments/e25-pi0-dppo) +完整跑通过这个组合。 ```bash plugrl-run-server pi0-policy default eval default \ @@ -36,8 +55,8 @@ plugrl-run-server pi0-policy default eval default \ ## 常见问题 -- CLI 找不到 UID:注册模块没有被 import。 -- OpenPI import 失败:按 server 仓库里的 OpenPI README 完成本地依赖。 +- CLI 找不到 UID:注册模块没有被 import,或者这个策略的依赖没装(见上文)。 +- `pi0-policy` 启动失败:按 server 仓库里的 OpenPI README 完成本地依赖。 ## 下一步