Skip to content
Closed
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
41 changes: 41 additions & 0 deletions docs/flashsac_optimization_report_zh.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
# FlashSAC 优化简报

## 改了什么

之前的 FlashSAC 优化把编译边界从单独的 loss 扩大到完整 objective,让 Inductor 可以跨越网络前向、categorical target 和 loss 进行融合;同时保留 graph-safe 的固定输出与延迟 metrics 拷贝,避免 CUDA Graph replay 覆盖正在使用的地址。

## 为什么有效

原路径会产生多个小 kernel 和中间 tensor。完整 objective compile 减少了 launch 和中间读写,固定输出则避免编译器临时 buffer 与 metrics buffer 复用同一地址。collector 仍然使用真实环境进程,learner 继续在 GPU 上更新。

## 怎么看真实训练耗时

```bash
uv run scripts/benchmark/rl/benchmark_flashsac_training.py \
--backend mujoco --iterations 20 --num-envs 256 \
--uni-rl-src ../unilab_rl-flashsac-optimization/src \
--output scripts/benchmark/outputs/flashsac_training/summary.json
```

脚本会显示每轮的:

```text
step learner(ms) collector(ms) iter(ms) reward
...
```

结尾输出三组 mean、median、p90、p95。这里的 collector 时间来自真实 MuJoCo/Motrix `env.step()`,不是预生成 action 的离线测试;learner 时间来自真实 FlashSAC update。

## 结果解释

总轮时间不是 learner 和 collector 的简单串行相加,因为 double-buffer runner 会让两者部分重叠。判断优化是否有效,应比较同一 backend、task、num_envs、batch、iterations 下的完整 `iter_ms`,同时观察 learner 和 collector 两列,避免只看其中一个阶段。

本机用 G1WalkFlat + MuJoCo 做了短跑验证(16 env,batch 16,1 update/step)。eager 路径的 5 个有效样本结果为:learner `6.572 / 6.219 / 7.338 / 7.520 ms`,collector `49.814 / 9.006 / 123.028 / 140.249 ms`,总轮 `7.446 / 7.048 / 8.324 / 8.545 ms`(依次为 mean / median / p90 / p95)。开启 full-objective compile 后,去掉前 4 个编译/预热样本,3 个 steady-state 样本为:learner `2.217 / 2.215 / 2.265 / 2.271 ms`,collector `4.110 / 3.914 / 4.529 / 4.606 ms`,总轮 `3.015 / 3.020 / 3.061 / 3.066 ms`。样本很少,这组数字用于验证计时链路,不作为正式跨配置结论。

## 复现注意事项

- 需要安装 UniLab 对应的物理 backend extra;MuJoCo 是最容易复现的选择。
- 测试前先确认 `unilab_rl` 版本,推荐使用 `--uni-rl-src` 指向包含 FlashSAC 优化提交的 checkout。
- 首轮可能包含 torch.compile 编译开销;比较 steady state 时应增加 iterations,并丢弃最初的编译/预热轮次。
- 脚本可用 `--skip-first N` 丢弃前 N 条 timing 样本,JSON 仍保留完整的 `rows_all`。
- 本报告不包含 Triton categorical-target 实验;Triton 仍是单独的可选实验目录。
55 changes: 55 additions & 0 deletions scripts/benchmark/rl/README_flashsac_training.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
# FlashSAC 真实训练计时

`benchmark_flashsac_training.py` 会调用生产入口 `train_flashsac.py`,创建真实物理后端环境,并实时转发训练输出。训练结束后,它从同一次训练的 TensorBoard 日志中打印:

- 每轮 learner update 时间;
- 每轮 collector cycle 时间(环境步进、replay 写入和 bookkeeping);
- 每轮总时间和 reward;
- learner、collector、总时间的 mean / median / p90 / p95。

## 运行

MuJoCo:

```bash
uv run scripts/benchmark/rl/benchmark_flashsac_training.py \
--backend mujoco --iterations 20 --num-envs 256 \
--output scripts/benchmark/outputs/flashsac_training/summary.json
```

Motrix:

```bash
uv run scripts/benchmark/rl/benchmark_flashsac_training.py \
--backend motrix --iterations 20 --num-envs 256
```

默认开启之前 FlashSAC 优化使用的 full-objective compile。要测 eager 路径,使用 `--no-compile`。如果要明确使用本地的优化版 `unilab_rl` checkout:

```bash
uv run scripts/benchmark/rl/benchmark_flashsac_training.py \
--uni-rl-src ../unilab_rl-flashsac-optimization/src \
--backend mujoco --iterations 20 --num-envs 256
```

compile 首轮可能包含编译开销;要只看 steady-state,可排除前几条共同 timing 样本:

```bash
... --skip-first 4
```

JSON 同时保留未过滤的 `rows_all` 和用于统计的 `rows`。

实际训练日志会保存在 `scripts/benchmark/outputs/flashsac_training/<timestamp>/`,可以直接用 TensorBoard 查看:

```bash
uv run tensorboard --logdir scripts/benchmark/outputs/flashsac_training
```

`--task` 可替换为其他 FlashSAC owner task,只要对应的 `<task>/<backend>.yaml` 存在。`--extra-override key=value` 可以追加 Hydra 参数,例如:

```bash
--extra-override training.torch_threads.learner_num_threads=8
```

注意:这不是模型或 collector 的仿真 microbenchmark,而是完整的 collector→replay→learner 训练链路;collector 时间来自真实物理后端的 `env.step()`。
279 changes: 279 additions & 0 deletions scripts/benchmark/rl/benchmark_flashsac_training.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,279 @@
#!/usr/bin/env python3
"""Run a real FlashSAC training loop and report learner/collector timing.

Unlike the collector-only benchmark, this command launches the production
``train_flashsac.py`` entrypoint, constructs the selected physics backend, and
reads the timing scalars emitted by the real off-policy runner. The child
process output is forwarded live so the user can see training progress.

Example (MuJoCo):

uv run scripts/benchmark/rl/benchmark_flashsac_training.py \
--backend mujoco --iterations 20 --num-envs 256

The same command works with ``--backend motrix`` when the Motrix extra is
installed. Use ``--uni-rl-src`` to test a local unilab-rl checkout instead of
the installed package, for example the FlashSAC optimization branch.
"""

from __future__ import annotations

import argparse
import json
import os
import subprocess
import sys
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from statistics import fmean
from typing import Iterable, Sequence

ROOT_DIR = Path(__file__).resolve().parents[3]
TRAIN_SCRIPT = ROOT_DIR / "src" / "unilab" / "scripts" / "train_flashsac.py"
DEFAULT_OUTPUT_ROOT = ROOT_DIR / "scripts" / "benchmark" / "outputs" / "flashsac_training"
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))

from scripts.benchmark.core.device_info import get_device_info_dict


@dataclass(frozen=True)
class Scalar:
step: int
value: float


def _percentile(values: Sequence[float], quantile: float) -> float:
if not values:
raise ValueError("cannot calculate a percentile of an empty sequence")
ordered = sorted(float(value) for value in values)
position = (len(ordered) - 1) * quantile
lower = int(position)
upper = min(lower + 1, len(ordered) - 1)
return ordered[lower] + (position - lower) * (ordered[upper] - ordered[lower])


def summarize(values: Iterable[float]) -> dict[str, float | int]:
samples = [float(value) for value in values]
if not samples:
raise ValueError("no timing samples were emitted")
return {
"count": len(samples),
"mean_ms": fmean(samples),
"median_ms": _percentile(samples, 0.50),
"p90_ms": _percentile(samples, 0.90),
"p95_ms": _percentile(samples, 0.95),
"min_ms": min(samples),
"max_ms": max(samples),
}


def _event_file(run_dir: Path) -> Path:
candidates = sorted(run_dir.rglob("events.out.tfevents.*"))
if not candidates:
raise RuntimeError(f"no TensorBoard event file found under {run_dir}")
return candidates[-1]


def _read_scalars(run_dir: Path, tag: str) -> list[Scalar]:
from tensorboard.backend.event_processing import event_accumulator

accumulator = event_accumulator.EventAccumulator(str(_event_file(run_dir)))
accumulator.Reload()
if tag not in accumulator.Tags().get("scalars", []):
return []
return [Scalar(int(event.step), float(event.value)) for event in accumulator.Scalars(tag)]


def _by_step(samples: Sequence[Scalar]) -> dict[int, float]:
return {sample.step: sample.value for sample in samples}


def parse_timing(run_dir: Path, *, skip_first: int = 0) -> dict[str, object]:
"""Read per-iteration learner, collector, wall and reward timings."""
learner = _read_scalars(run_dir, "timing/learner_train_ms")
collector = _read_scalars(run_dir, "perf/collector_cycle_ms")
wall = _read_scalars(run_dir, "perf/iter_ms")
reward = _read_scalars(run_dir, "reward/mean")
if not learner:
raise RuntimeError("training log has no timing/learner_train_ms samples")
if not collector:
# Older runners do not have the aggregate scalar, but emit the three
# mutually exclusive collector phases. Reconstruct the cycle exactly.
phases = [
_by_step(_read_scalars(run_dir, f"timing/collector_{name}"))
for name in ("env_step_ms", "replay_ms", "bookkeeping_ms")
]
steps = sorted(set().union(*(phase.keys() for phase in phases)))
collector = [Scalar(step, sum(phase.get(step, 0.0) for phase in phases)) for step in steps]
if not collector:
raise RuntimeError("training log has no collector timing samples")

learner_by_step = _by_step(learner)
collector_by_step = _by_step(collector)
wall_by_step = _by_step(wall)
reward_by_step = _by_step(reward)
steps = sorted(set(learner_by_step) & set(collector_by_step))
rows_all = [
{
"step": step,
"learner_train_ms": learner_by_step[step],
"collector_cycle_ms": collector_by_step[step],
"iter_ms": wall_by_step.get(step),
"reward": reward_by_step.get(step),
}
for step in steps
]
if skip_first < 0:
raise ValueError("skip_first must be non-negative")
rows = rows_all[skip_first:]
if not rows:
raise RuntimeError(
f"learner and collector timing series have fewer than {skip_first + 1} common steps"
)
return {
"num_samples": len(rows),
"num_samples_all": len(rows_all),
"skip_first": skip_first,
"rows_all": rows_all,
"learner_train_ms": summarize(row["learner_train_ms"] for row in rows),
"collector_cycle_ms": summarize(row["collector_cycle_ms"] for row in rows),
"iter_ms": summarize(row["iter_ms"] for row in rows if row["iter_ms"] is not None),
"rows": rows,
}


def _command(args: argparse.Namespace, run_dir: Path) -> list[str]:
task = f"{args.task}/{args.backend}"
command = [
sys.executable,
str(TRAIN_SCRIPT),
f"task={task}",
"training.no_play=true",
f"training.sim_backend={args.backend}",
f"training.log_dir={run_dir}",
f"algo.max_iterations={args.iterations}",
f"algo.num_envs={args.num_envs}",
f"algo.batch_size={args.batch_size}",
f"algo.replay_buffer_n={args.replay_buffer_n}",
f"algo.learning_starts={args.learning_starts}",
f"algo.updates_per_step={args.updates_per_step}",
"algo.save_interval=1000000",
f"algo.algo_params.use_compile={str(args.compile).lower()}",
f"algo.algo_params.compile_full_objectives={str(args.compile).lower()}",
]
command.extend(args.extra_override)
return command


def _run(command: Sequence[str], *, env: dict[str, str]) -> None:
print("$ " + " ".join(command), flush=True)
process = subprocess.Popen(
list(command),
cwd=ROOT_DIR,
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
)
assert process.stdout is not None
for line in process.stdout:
print(line, end="", flush=True)
return_code = process.wait()
if return_code != 0:
raise subprocess.CalledProcessError(return_code, list(command))


def _print_report(report: dict[str, object]) -> None:
rows = report["rows"]
assert isinstance(rows, list)
print("\nPer-iteration timing (ms):")
print(f"{'step':>8} {'learner':>12} {'collector':>12} {'iter':>12} {'reward':>12}")
for row in rows:
assert isinstance(row, dict)
reward = row.get("reward")
reward_text = f"{float(reward):12.4f}" if reward is not None else f"{'-':>12}"
iter_time = row.get("iter_ms")
iter_text = f"{float(iter_time):12.3f}" if iter_time is not None else f"{'-':>12}"
print(
f"{int(row['step']):8d} {float(row['learner_train_ms']):12.3f} "
f"{float(row['collector_cycle_ms']):12.3f} {iter_text} {reward_text}"
)
print("\nSummary (ms):")
print(f"{'phase':<18} {'mean':>10} {'median':>10} {'p90':>10} {'p95':>10} {'n':>6}")
for name in ("learner_train_ms", "collector_cycle_ms", "iter_ms"):
stats = report[name]
assert isinstance(stats, dict)
print(
f"{name:<18} {float(stats['mean_ms']):10.3f} {float(stats['median_ms']):10.3f} "
f"{float(stats['p90_ms']):10.3f} {float(stats['p95_ms']):10.3f} {int(stats['count']):6d}"
)


def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--backend", choices=("mujoco", "motrix"), default="mujoco")
parser.add_argument("--task", default="g1_walk_flat")
parser.add_argument("--iterations", type=int, default=20)
parser.add_argument("--num-envs", type=int, default=256)
parser.add_argument("--batch-size", type=int, default=256)
parser.add_argument("--replay-buffer-n", type=int, default=32)
parser.add_argument("--learning-starts", type=int, default=8)
parser.add_argument("--updates-per-step", type=int, default=2)
parser.add_argument("--compile", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--run-dir", type=Path, default=None)
parser.add_argument("--output", type=Path, default=None)
parser.add_argument("--keep-run", action="store_true")
parser.add_argument(
"--skip-first",
type=int,
default=0,
help="Exclude the first N common timing rows from summary statistics (compile warm-up).",
)
parser.add_argument("--uni-rl-src", type=Path, default=None)
parser.add_argument("--extra-override", action="append", default=[])
return parser.parse_args(argv)


def main(argv: Sequence[str] | None = None) -> int:
args = parse_args(argv)
if args.run_dir is None:
run_dir = DEFAULT_OUTPUT_ROOT / datetime.now().strftime("%Y%m%d_%H%M%S")
else:
run_dir = args.run_dir.resolve()
run_dir.mkdir(parents=True, exist_ok=True)
env = os.environ.copy()
source_paths = [str(ROOT_DIR / "src")]
if args.uni_rl_src is not None:
source_paths.insert(0, str(args.uni_rl_src.resolve()))
elif (ROOT_DIR.parent / "unilab_rl-flashsac-optimization" / "src").is_dir():
source_paths.insert(0, str(ROOT_DIR.parent / "unilab_rl-flashsac-optimization" / "src"))
env["PYTHONPATH"] = os.pathsep.join(source_paths + [env.get("PYTHONPATH", "")])
command = _command(args, run_dir)
_run(command, env=env)
report = parse_timing(run_dir, skip_first=args.skip_first)
payload = {
"command": command,
"run_dir": str(run_dir),
"backend": args.backend,
"task": args.task,
"device": get_device_info_dict(),
**report,
}
_print_report(report)
if args.output is not None:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
print(f"\nSaved JSON: {args.output}")
if not args.keep_run and args.output is None:
# Keep the run by default when --output is requested; otherwise leave
# it available for TensorBoard inspection because it contains the
# actual physics-backed training artifacts.
print(f"Training logs: {run_dir}")
return 0


if __name__ == "__main__":
raise SystemExit(main())
7 changes: 6 additions & 1 deletion src/unilab/base/np_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -614,7 +614,12 @@ def play_capabilities(self) -> EnvPlayCapabilities:
supports_native_video_capture=capabilities.supports_native_video_capture,
supports_debug_overlay=capabilities.supports_debug_overlay,
supports_interactive_debug_overlay=capabilities.supports_interactive_debug_overlay,
supports_mocap_playback=capabilities.supports_mocap_playback,
# Older unisim-core releases do not expose mocap playback on the
# backend capability record. Keep the env contract fail-closed so
# collectors can still run against those physics adapters.
supports_mocap_playback=bool(
getattr(capabilities, "supports_mocap_playback", False)
),
)

def get_playback_model(self, env_index: int | None = None) -> Any:
Expand Down
Loading
Loading