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
32 changes: 32 additions & 0 deletions .github/workflows/unit-tests.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
name: Unit Tests

on:
push:
branches:
- main
- master
pull_request:

jobs:
unit-tests:
runs-on: ubuntu-latest

steps:
- name: Checkout
uses: actions/checkout@v4

- name: Setup Python
uses: actions/setup-python@v5
with:
python-version: "3.11"

- name: Setup uv
uses: astral-sh/setup-uv@v5

- name: Install project and test dependencies
run: |
uv venv .venv
uv pip install --python .venv/bin/python --torch-backend cpu -e '.[dev]'

- name: Run unit tests
run: .venv/bin/python -m pytest tests/unit
4 changes: 4 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ pip-delete-this-directory.txt
htmlcov/
.tox/
.nox/
.pytest_artifacts/
.coverage
.coverage.*
.cache
Expand Down Expand Up @@ -140,6 +141,9 @@ logs/
runs/
outputs/
output/
.o/
.rt/
.w/
# runs
resource_pool_auto.yaml

Expand Down
9 changes: 9 additions & 0 deletions examples/wbc_tracking/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,18 @@ override-dependencies = [
"sympy==1.13.1"
]

[[tool.uv.index]]
name = "pypi"
url = "https://pypi.org/simple"

[[tool.uv.index]]
name = "nvidia"
url = "https://pypi.nvidia.com"

[tool.uv.sources]
rlightning = {path = "../../", editable = true}
rsl-rl = {path = "../../third_party/rsl_rl", editable = true}
isaaclab = { index = "nvidia" }

[tool.uv.extra-build-dependencies]
flatdict = ["setuptools<81"]
29 changes: 26 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,3 +1,7 @@
[build-system]
requires = ["setuptools>=80.8.0,<81", "wheel"]
build-backend = "setuptools.build_meta"

[project]
name = "rlightning"
version = "0.1.0"
Expand Down Expand Up @@ -26,14 +30,16 @@ dependencies = [
"transformers",
"uvloop",
"wandb",
"zmq",
"pyzmq",
"zstandard>=0.25.0",
]

[project.optional-dependencies]
dev = [
"debugpy>=1.8.14",
"ipython",
"pytest>=8.3.5",
"pytest-cov>=6.1.1",
]
isaaclab = [
"isaaclab[isaacsim]==2.2.0",
Expand All @@ -54,13 +60,17 @@ humanoid = [
"scipy",
"easydict",
"joblib",
"smplx @ git+https://github.com/ZhengyiLuo/smplx.git@master",
"smpl-sim @ git+https://github.com/ZhengyiLuo/SMPLSim.git@master",
"open3d==0.19.0",
"natsort==8.4.0",
"mink==0.0.13",
]

[dependency-groups]
humanoid-dev = [
"smplx @ git+https://github.com/ZhengyiLuo/smplx.git@master",
"smpl-sim @ git+https://github.com/ZhengyiLuo/SMPLSim.git@master",
]


[tool.isort]
profile = "black"
Expand All @@ -73,6 +83,19 @@ typeCheckingMode = "off"
reportMissingImports = false
reportMissingModuleSource = false

[tool.pytest.ini_options]
minversion = "8.0"
testpaths = ["tests"]
addopts = "-ra -m 'not integration and not slow and not gpu and not maniskill and not isaaclab and not e2e'"
markers = [
"slow: marks tests that are slower or intended for scheduled regression runs",
"gpu: marks tests that require a GPU runtime",
"integration: marks tests that cover multiple components working together",
"e2e: marks end-to-end CLI or workflow tests",
"maniskill: marks tests that require the ManiSkill stack",
"isaaclab: marks tests that require the IsaacLab stack",
]

[tool.setuptools.packages.find]
include = ["rlightning*"]

Expand Down
2 changes: 1 addition & 1 deletion rlightning/engine/async_rl_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -249,7 +249,7 @@ def _train_loop(self) -> None:
self.coordinator.wait_for_dataset_ready()
self.coordinator.wait_for_weights_updated()
self._train()
if self.config.train.save_interval > 0 and (self.epoch + 1) % self.config.train.save_interval == 0:
if self.config.train.save_interval > 0 and self.epoch % self.config.train.save_interval == 0:
ckpt_path = f"{self.config.train.save_dir}/epoch_{self.epoch}.pt"
self.policy_group.save_checkpoint(path=ckpt_path)
self.coordinator.notify_train_step_done()
Expand Down
12 changes: 5 additions & 7 deletions rlightning/engine/sync_rl_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -215,21 +215,19 @@ def run(self) -> None:
Executes the training loop for the configured number of epochs,
performing rollout, training, and periodic evaluation.
"""
logger.info("Evaluating before training...")
self._evaluate(obj_set="train", prefix="eval")
self._evaluate(obj_set="test", prefix="eval_ood")

for self.epoch in self.iter_epochs(num_epochs=self.config.train.max_epochs):
self._rollout(obj_set="train", prefix="rollout")
self._update_dataset()
self._train()
self._sync_weights()

if self.config.train.eval_interval > 0 and (self.epoch + 1) % self.config.train.save_interval == 0:
if self.config.train.eval_interval > 0 and self.epoch % self.config.train.save_interval == 0:
ckpt_path = f"{self.config.train.save_dir}/epoch_{self.epoch}.pt"
self.policy_group.save_checkpoint(path=ckpt_path)

# sync weights after training and save checkpoint
self._sync_weights()

if self.config.train.eval_interval > 0 and (self.epoch + 1) % self.config.train.eval_interval == 0:
if self.config.train.eval_interval > 0 and self.epoch % self.config.train.eval_interval == 0:
logger.info(f"Evaluating at epoch {self.epoch}")
self._evaluate(obj_set="train", prefix="eval")
self._evaluate(obj_set="test", prefix="eval_ood")
Expand Down
38 changes: 30 additions & 8 deletions rlightning/policy/base_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,23 @@ def infer_train_dp_world_size() -> int:
return 1


def clone_checkpoint_value(value: Any) -> Any:
"""Recursively clone checkpoint payloads into CPU-owned tensors."""
if torch.is_tensor(value):
return value.detach().cpu().clone()

if isinstance(value, dict):
return value.__class__((key, clone_checkpoint_value(child)) for key, child in value.items())

if isinstance(value, list):
return [clone_checkpoint_value(child) for child in value]

if isinstance(value, tuple):
return tuple(clone_checkpoint_value(child) for child in value)

return value


class PolicyRole(StrEnum):
"""Policy role enumeration."""

Expand Down Expand Up @@ -548,15 +565,20 @@ def save_checkpoint(self, path: str) -> None:
ckpt_folder = Path(path).parent
os.makedirs(ckpt_folder, exist_ok=True)

state: Dict[str, Dict] = {}
for name, model in self.model_list:
if isinstance(model, DDP):
module = model.module
else:
module = model
state[name] = module.state_dict()
model_was_offloaded = getattr(self, "_model_params_offloaded", False)
if model_was_offloaded:
self.reload_model_param_and_grad(load_grad=False)

torch.save(state, path)
try:
state: Dict[str, Dict] = {}
for name, model in self.model_list:
module = model.module if isinstance(model, DDP) else model
state[name] = clone_checkpoint_value(module.state_dict())

torch.save(state, path)
finally:
if model_was_offloaded:
self.offload_model_param_and_grad(offload_grad=False)

def reset_training_state(
self, train_config: TrainConfig, env_meta: Optional[Any] = None, seed: Optional[int] = None
Expand Down
5 changes: 4 additions & 1 deletion rlightning/weights/weight_buffer_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ def __init_weight_buffer_mixin__(self, buffer_strategy: str):

# for offload model param and grad
self.cpu_param_backup = {}
self._model_params_offloaded = False

def init_weight_buffer(self, shared_weight_buffer=None):
"""Initialize the weight buffer."""
Expand Down Expand Up @@ -177,10 +178,11 @@ def offload_model_param_and_grad(self, offload_grad=False):
actual_model = self.model.module if isinstance(self.model, DDP) else self.model
for name, param in actual_model.named_parameters():
if param.data.storage().size() > 0:
self.cpu_param_backup[name] = (param.data.cpu(), param.data.size())
self.cpu_param_backup[name] = (param.data.detach().cpu().clone(), param.data.size())
_free_storage(param.data)
if offload_grad and param.grad is not None:
param.grad = param.grad.to("cpu", non_blocking=True)
self._model_params_offloaded = True
self.clear_memory(sync=True)
profiler.log_gpu_memory_usage("offload_model_param_and_grad")

Expand All @@ -198,6 +200,7 @@ def reload_model_param_and_grad(self, load_grad=False):

if load_grad and param.grad is not None:
param.grad = param.grad.to(self.device, non_blocking=True)
self._model_params_offloaded = False
self.clear_memory(sync=True)
profiler.log_gpu_memory_usage("reload_model_param_and_grad")

Expand Down
Loading
Loading