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
2 changes: 1 addition & 1 deletion CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ Languages: English | [简体中文](docs/sphinx/source/zh_CN/4-developer_guide/4
1. Fork and clone the repository.
2. Install dependencies for your platform:
- macOS (MPS, installs PyPI torch wheels): `make setup-motrix` (or `make setup-mujoco`)
- Linux default (installs PyTorch cu128 wheels; requires an NVIDIA GPU/driver supported by current PyTorch cu128 wheels): `make setup`
- Linux default (installs PyTorch cu130 wheels; requires an NVIDIA GPU/driver supported by current PyTorch cu130 wheels): `make setup`
- Linux AMD / ROCm workstation: `make sync-rocm`, then run commands with `uv run --no-sync ...`
- For direct uv setup, use `uv sync --extra mujoco --extra motrix`; replace it
with `--extra mujoco` or `--extra motrix` for a single backend
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -245,7 +245,7 @@ Reverse-lookup from error text to cause and fix.
## Platform Profiles

Linux CUDA and macOS use the default `pyproject.toml`. The default Linux torch
wheel source is the PyTorch `cu128` index configured in `pyproject.toml`.
wheel source is the PyTorch `cu130` index configured in `pyproject.toml`.

On Apple Silicon macOS, `make setup-motrix` is the shortest interactive path.
The CLI routes Motrix playback through `mxpython` when needed; MuJoCo playback
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ Install dependencies for your platform. The setup targets also install the
optional simulator extras used by the repository's checks:

- macOS (MPS, PyPI torch wheel): `make setup-motrix` (or `uv sync --extra mujoco`)
- Linux with NVIDIA (PyTorch cu128 wheel): `make setup`
- Linux with NVIDIA (PyTorch cu130 wheel): `make setup`
- Linux AMD / ROCm: `make sync-rocm`, then run commands with `uv run ...`. To
return to the default CUDA / macOS profile, `git restore -- pyproject.toml
uv.lock` and re-run `make setup`.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -220,7 +220,7 @@ fork 的构建在编译期钉住 `mujoco==3.11.0`,因此隔离构建总是针
## 平台配置档

Linux CUDA 和 macOS 使用默认的 `pyproject.toml`。默认的 Linux torch
wheel 来源是在 `pyproject.toml` 中配置的 PyTorch `cu128` 索引。
wheel 来源是在 `pyproject.toml` 中配置的 PyTorch `cu130` 索引。

在 Apple Silicon macOS 上,`make setup-motrix` 是最短的交互式路径。CLI 会在需要时
通过 `mxpython` 路由 Motrix 回放;MuJoCo 回放使用官方 MuJoCo wheel 自带的
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
按平台安装依赖。setup target 也会安装仓库检查所需的可选仿真器 extra:

- macOS(MPS,PyPI torch wheel):`make setup-motrix`(或 `uv sync --extra mujoco`)
- Linux NVIDIA(PyTorch cu128 wheel):`make setup`
- Linux NVIDIA(PyTorch cu130 wheel):`make setup`
- Linux AMD / ROCm:`make sync-rocm`,随后用 `uv run ...` 运行命令。要切回默认
CUDA / macOS profile,执行 `git restore -- pyproject.toml uv.lock` 后重新
`make setup`。
Expand Down
6 changes: 3 additions & 3 deletions pyproject.rocm.toml
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,8 @@ requires-python = ">=3.10,<3.14"
dependencies = [
"numpy",
"unisim-core>=1.7.8",
"torch==2.11.0",
"triton-rocm==3.6.0 ; sys_platform == 'linux' and platform_machine == 'x86_64'",
"torch==2.14.0",
"triton-rocm==3.8.0 ; sys_platform == 'linux' and platform_machine == 'x86_64'",
"gymnasium",
"imageio",
"etils",
Expand All @@ -35,7 +35,7 @@ dependencies = [
"packaging",
"mediapy",
"tensorboard",
"setuptools<70",
"setuptools",
"rich",
"tqdm",
"typing-extensions",
Expand Down
13 changes: 3 additions & 10 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -40,8 +40,7 @@ dependencies = [
"prettytable>=3.10",
# A range (not an exact pin) lets ROCm users substitute a ROCm torch
# build; uv.lock pins the tested CUDA builds via tool.uv.sources.
"torch>=2.9,<2.12 ; sys_platform == 'linux' and platform_machine == 'aarch64'",
"torch>=2.8,<2.12 ; sys_platform != 'linux' or platform_machine != 'aarch64'",
"torch>=2.9,<2.15",
"gymnasium",
"imageio",
"etils",
Expand Down Expand Up @@ -122,21 +121,15 @@ superdex = [
# unilab-rl runtime stays installed in the dev environment.
dev = ["pytest", "pytest-cov", "ruff", "mypy", "pyright>=1.1.408", "unilab-rl==1.4.0"]

[[tool.uv.index]]
name = "pytorch-cu128"
url = "https://download.pytorch.org/whl/cu128"
explicit = true

[[tool.uv.index]]
name = "r2-cu130"
url = "https://download-r2.pytorch.org/whl/cu130"
explicit = true

[tool.uv.sources]
# cu130 carries torch>=2.12 CUDA wheels for linux/win; macOS resolves from PyPI.
torch = [
{ index = "r2-cu130", marker = "sys_platform=='linux' and platform_machine=='aarch64'" },
{ index = "pytorch-cu128", marker = "sys_platform=='linux' and platform_machine=='x86_64'" },
{ index = "pytorch-cu128", marker = "sys_platform=='win32'" },
{ index = "r2-cu130", marker = "sys_platform=='linux' or sys_platform=='win32'" },
]

[tool.uv]
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/conf/appo/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,9 @@ training:
device: null
collector_device: null
logger: tensorboard
# Backend (TensorBoard/wandb) scalar reporting interval in iterations; the
# terminal display is unaffected. Each report is one batched event record.
log_interval: 1
wandb_project: unilab
wandb_entity: null
wandb_group: null
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/conf/flashsac/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,9 @@ training:
# MuJoCo pool; null = auto partition of cpu_count // world_size per rank.
dp_collector_cpu_ids: null
logger: tensorboard
# Backend (TensorBoard/wandb) scalar reporting interval in iterations; the
# terminal display is unaffected. Each report is one batched event record.
log_interval: 1
wandb_project: unilab
wandb_entity: null
wandb_group: null
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/conf/ppo/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,9 @@ training:
dp_collector_cpu_ids: null
device: null
logger: tensorboard
# Backend (TensorBoard/wandb) scalar reporting interval in iterations; the
# terminal display is unaffected. Each report is one batched event record.
log_interval: 1
wandb_project: unilab
wandb_entity: null
wandb_group: null
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/conf/sac/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,9 @@ training:
# MuJoCo pool; null = auto partition of cpu_count // world_size per rank.
dp_collector_cpu_ids: null
logger: tensorboard
# Backend (TensorBoard/wandb) scalar reporting interval in iterations; the
# terminal display is unaffected. Each report is one batched event record.
log_interval: 1
wandb_project: unilab
wandb_entity: null
wandb_group: null
Expand Down
3 changes: 3 additions & 0 deletions src/unilab/conf/warpsac/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,9 @@ training:
devices: null
dp_collector_cpu_ids: null
logger: tensorboard
# Backend (TensorBoard/wandb) scalar reporting interval in iterations; the
# terminal display is unaffected. Each report is one batched event record.
log_interval: 1
wandb_project: unilab
wandb_entity: null
wandb_group: null
Expand Down
1 change: 1 addition & 0 deletions src/unilab/scripts/train_appo.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@ def build_appo_runner_kwargs(
"steps_per_env": cfg.algo.steps_per_env,
"sim_backend": cfg.training.sim_backend,
"seed": rl_cfg.get("seed"),
"log_interval": int(cfg.training.log_interval),
}
if cfg.training.replay_queue_size is not None:
runner_kwargs["replay_queue_size"] = cfg.training.replay_queue_size
Expand Down
4 changes: 4 additions & 0 deletions src/unilab/scripts/train_rsl_rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
ExperimentTracker,
patch_rsl_rl_action_std_logging,
patch_rsl_rl_resume_state,
patch_rsl_rl_tensorboard_logging,
patch_rsl_rl_wandb_writer,
)
from unilab.utils.checkpoint import get_entrypoint_log_root
Expand Down Expand Up @@ -615,6 +616,9 @@ def main(cfg: DictConfig) -> None:
),
)
patch_rsl_rl_action_std_logging(runner)
patch_rsl_rl_tensorboard_logging(
runner, log_interval=int(cfg.training.log_interval)
)

if cfg.algo.load_run != "-1":
resume_path, _ = parse_checkpoint_path(cfg, root_dir=Path.cwd())
Expand Down
122 changes: 122 additions & 0 deletions src/unilab/training/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,18 @@ def _load_wandb() -> Any | None:
return None


def _tb_proto_modules() -> tuple[Any, Any]:
"""Import tensorboard's protobuf modules as ``Any``.

The generated modules define their messages through reflection, which
static checkers cannot see; returning them as ``Any`` keeps attribute
access clean without per-line suppressions.
"""
from tensorboard.compat.proto import event_pb2, summary_pb2

return event_pb2, summary_pb2


def _json_safe(value: Any) -> Any:
if isinstance(value, dict):
return {str(k): _json_safe(v) for k, v in value.items()}
Expand Down Expand Up @@ -376,6 +388,116 @@ def _safe_log(self: Any, *args: Any, **kwargs: Any) -> Any:
runner.logger.log = _safe_log.__get__(runner.logger, type(runner.logger))


class _BatchedScalarWriter:
"""Buffer rsl-rl's per-tag ``add_scalar`` calls and emit one event per step.

TensorBoard's writer thread performs one open/write/close per record, which
saturates its async queue (depth 10) and blocks the training loop when the
log directory lives on a network filesystem. Batching one iteration's
scalars into a single ``Event`` keeps it to one record per step group.
"""

def __init__(self, writer: Any) -> None:
self._writer = writer
self._pending: list[tuple[str, float]] = []
self._pending_step: int | None = None

def add_scalar(self, tag: str, value: Any, step: Any = None, *args: Any, **kwargs: Any) -> None:
try:
scalar = float(value)
except (TypeError, ValueError):
self._flush()
self._writer.add_scalar(tag, value, step, *args, **kwargs)
return
step_key = int(step) if step is not None else 0
if self._pending and step_key != self._pending_step:
self._flush()
self._pending_step = step_key
self._pending.append((tag, scalar))

def _flush(self) -> None:
if not self._pending:
return
event_pb2, summary_pb2 = _tb_proto_modules()
event = event_pb2.Event(
wall_time=time.time(),
step=self._pending_step or 0,
summary=summary_pb2.Summary(
value=[
summary_pb2.Summary.Value(tag=tag, simple_value=value)
for tag, value in self._pending
]
),
)
self._writer.file_writer.add_event(event)
self._pending = []
self._pending_step = None

def flush(self) -> None:
self._flush()
self._writer.flush()

def close(self) -> None:
self._flush()
self._writer.close()

def __getattr__(self, name: str) -> Any:
return getattr(self._writer, name)


class _SwallowScalarWriter:
"""Drop ``add_scalar`` writes while staying truthy for rsl-rl's console log."""

def __init__(self, writer: Any) -> None:
self._writer = writer

def add_scalar(self, *args: Any, **kwargs: Any) -> None:
pass

def __getattr__(self, name: str) -> Any:
return getattr(self._writer, name)


def patch_rsl_rl_tensorboard_logging(runner: Any, log_interval: int = 1) -> None:
"""Batch and optionally throttle rsl-rl's TensorBoard scalar writes.

Wraps ``runner.logger.writer`` so each step's ~15-30 ``add_scalar`` calls
become a single event record, and gates backend writes to every
``log_interval`` iterations (console output and episode bookkeeping are
unaffected; the final iteration is always logged). No-op for non-TB
backends and when no writer exists.
"""
logger = getattr(runner, "logger", None)
writer = getattr(logger, "writer", None)
if logger is None or writer is None:
return
if getattr(logger, "logger_type", "tensorboard") != "tensorboard":
return
if isinstance(writer, _BatchedScalarWriter):
return

batched = _BatchedScalarWriter(writer)
logger.writer = batched

interval = max(1, int(log_interval))
if interval <= 1:
return

swallow = _SwallowScalarWriter(batched)
original_log = logger.log

def _gated_log(self: Any, it: int, *args: Any, **kwargs: Any) -> Any:
total_it = args[1] if len(args) > 1 else int(kwargs.get("total_it", 0))
due = it % interval == 0 or it >= total_it
logger.writer = batched if due else swallow
try:
return original_log(it, *args, **kwargs)
finally:
logger.writer = batched

logger.log = _gated_log.__get__(logger, type(logger))


def patch_rsl_rl_wandb_writer() -> None:
"""Patch rsl-rl W&B writer so it can reuse an already-open run."""
try:
Expand Down
1 change: 1 addition & 0 deletions tests/config/test_config_system.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,7 @@ def test_algo_config_composes(algo_dir: str, config_name: str):
cfg = _compose(algo_dir, config_name)
assert cfg.training.task_name
assert cfg.training.sim_backend == "mujoco"
assert int(cfg.training.log_interval) >= 1


def test_backend_task_files_keep_full_identity():
Expand Down
40 changes: 15 additions & 25 deletions tests/scripts/test_torch_cuda_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,15 +22,12 @@ def test_torch_cuda_source_covers_windows_and_linux() -> None:
}
return

cu128_sources = [source for source in torch_sources if source.get("index") == "pytorch-cu128"]
cu130_sources = [source for source in torch_sources if source.get("index") == "r2-cu130"]

assert {source["marker"] for source in cu128_sources} == {
"sys_platform=='linux' and platform_machine=='x86_64'",
"sys_platform=='win32'",
}
# cu128 tops out at torch 2.11; torch>=2.12 CUDA wheels for linux/win all
# ship from cu130, so a single source entry covers both.
assert [source["marker"] for source in cu130_sources] == [
"sys_platform=='linux' and platform_machine=='aarch64'"
"sys_platform=='linux' or sys_platform=='win32'"
]


Expand All @@ -48,42 +45,35 @@ def test_windows_lock_uses_cuda_torch() -> None:
if rocm_dependency is not None:
assert {
"name": "torch",
"version": "2.11.0+rocm7.2",
"version": "2.14.0+rocm7.2",
"source": {"registry": "https://download.pytorch.org/whl/rocm7.2"},
"marker": "platform_machine == 'x86_64' and sys_platform == 'linux'",
} in torch_dependencies
return

cu128_dependency = next(
cu130_dependency = next(
dep
for dep in torch_dependencies
if dep["version"] == "2.8.0+cu128"
and dep["source"] == {"registry": "https://download.pytorch.org/whl/cu128"}
if dep["version"] == "2.14.0+cu130"
and dep["source"] == {"registry": "https://download-r2.pytorch.org/whl/cu130"}
)
# uv adds impossible-extra guards for the Newton/mjwarp and Newton/mujoco
# conflict matrix. Assert the platform clauses remain present while
# allowing those generated guards to evolve.
assert "platform_machine == 'x86_64' and sys_platform == 'linux'" in cu128_dependency["marker"]
assert "sys_platform == 'win32'" in cu128_dependency["marker"]
cu130_dependency = next(
dep
for dep in torch_dependencies
if dep["version"] == "2.9.0+cu130"
and dep["source"] == {"registry": "https://download-r2.pytorch.org/whl/cu130"}
)
assert "platform_machine == 'aarch64' and sys_platform == 'linux'" in cu130_dependency["marker"]
assert "sys_platform == 'linux'" in cu130_dependency["marker"]
assert "sys_platform == 'win32'" in cu130_dependency["marker"]

torch_packages = [package for package in lock["package"] if package["name"] == "torch"]
cu128_package = next(
cu130_package = next(
package
for package in torch_packages
if package["source"] == {"registry": "https://download.pytorch.org/whl/cu128"}
if package["source"] == {"registry": "https://download-r2.pytorch.org/whl/cu130"}
)

assert cu128_package["version"] == "2.8.0+cu128"
assert cu130_package["version"] == "2.14.0+cu130"
assert any(
"sys_platform == 'win32'" in marker for marker in cu128_package["resolution-markers"]
"sys_platform == 'win32'" in marker for marker in cu130_package["resolution-markers"]
)

wheel_urls = [wheel["url"] for wheel in cu128_package["wheels"]]
assert any("torch-2.8.0%2Bcu128" in url and "win_amd64.whl" in url for url in wheel_urls)
wheel_urls = [wheel["url"] for wheel in cu130_package["wheels"]]
assert any("torch-2.14.0%2Bcu130" in url and "win_amd64.whl" in url for url in wheel_urls)
Loading
Loading