Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 84 additions & 25 deletions docs/algorithm/custom_algorithm.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,47 +6,91 @@ 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/<algo_uid>/`.
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

Put the code in one of these layouts.

### Built-in in `plugrl-server`

- `plugrl_server/algorithm/<algo_uid>/<algo_uid>.py` implementation
- `plugrl_server/algorithm/<algo_uid>/__init__.py` registration import
- `plugrl_server/algorithm/__init__.py` imports your package
- `plugrl_server/algorithm/<algo_uid>/<algo_uid>_config.py`: the config
dataclass, registered with `@register_algo_config`
- `plugrl_server/algorithm/<algo_uid>/<algo_uid>.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 <algo_uid> 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.

Expand All @@ -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"
Expand All @@ -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

Expand All @@ -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, {}

Expand All @@ -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:
Expand All @@ -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`.
Expand All @@ -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.
Expand All @@ -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 <policy> <variant> --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.
Expand Down
103 changes: 77 additions & 26 deletions docs/algorithm/custom_algorithm.zh.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,46 +6,82 @@

1. 在 `plugrl-server/src/plugrl_server/algorithm/<algo_uid>/` 下新建包。
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/<algo_uid>/<algo_uid>.py` 实现
- `plugrl_server/algorithm/<algo_uid>/__init__.py` 注册导入
- `plugrl_server/algorithm/__init__.py` 导入你的包
- `plugrl_server/algorithm/<algo_uid>/<algo_uid>_config.py`:配置 dataclass,用
`@register_algo_config` 注册
- `plugrl_server/algorithm/<algo_uid>/<algo_uid>.py`:算法类,用 `@register_algo` 注册
- `plugrl_server/algorithm/__init__.py`:把这两个模块都 import 进来

两个都要 import。只 import 配置模块的话,UID 会出现在 CLI 里,但运行时打印完配置
就报 `KeyError: 'Algorithm <algo_uid> 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`,不会被悄悄当成改名放过。

## 最小模板

Expand All @@ -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"
Expand All @@ -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

Expand All @@ -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, {}

Expand All @@ -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:
Expand All @@ -121,6 +167,7 @@ class YourAlgorithm(BaseAlgorithm):
## 设计规则

- 模型结构与动作生成参数放在 policy config。
- 运行长度用 `BaseAlgoConfig.global_steps`,不要另起一个字段。进度显示读的就是它。
- `learn()` 用外部数据时,显式维护环境步数与更新步数。
- 训练调度相关计数写进 `Checkpoint.meta`,并在 `load_checkpoint` 恢复。
- 不写分布式钩子就用 `BaseAlgorithm`,不要直接上 `DDPAlgorithm`。
Expand All @@ -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。
Expand All @@ -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 <policy> <variant> --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`,或补齐分布式钩子。
Expand Down
Loading
Loading