From 417410331647b8b962de10a79f840895ce1f3cea Mon Sep 17 00:00:00 2001 From: Sergey Arkhangelskiy Date: Sun, 17 May 2026 10:37:18 +0300 Subject: [PATCH 1/3] Fix OpenPI LR schedule to span full run and log LR to wandb --- positronic/vendors/openpi/_launch.py | 50 ++++++++++++++++++++++++++++ positronic/vendors/openpi/train.py | 3 +- 2 files changed, 52 insertions(+), 1 deletion(-) create mode 100644 positronic/vendors/openpi/_launch.py diff --git a/positronic/vendors/openpi/_launch.py b/positronic/vendors/openpi/_launch.py new file mode 100644 index 000000000..94276933e --- /dev/null +++ b/positronic/vendors/openpi/_launch.py @@ -0,0 +1,50 @@ +import dataclasses +import importlib.util +import sys + +import openpi.training.config as _config +import wandb +from openpi.training.optimizer import CosineDecaySchedule + + +def _extract_openpi_root() -> str: + argv = sys.argv[1:] + if len(argv) < 2 or argv[0] != '--openpi-root': + raise SystemExit('_launch.py expects `--openpi-root ` as its first argument') + root = argv[1] + sys.argv = [sys.argv[0], *argv[2:]] + return root + + +def _load_openpi_train(openpi_root: str): + spec = importlib.util.spec_from_file_location('openpi_train', f'{openpi_root}/scripts/train.py') + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def main(): + openpi_root = _extract_openpi_root() + openpi_train = _load_openpi_train(openpi_root) + + cfg = _config.cli() + + if isinstance(cfg.lr_schedule, CosineDecaySchedule): + cfg = dataclasses.replace( + cfg, lr_schedule=dataclasses.replace(cfg.lr_schedule, decay_steps=cfg.num_train_steps) + ) + + lr_fn = cfg.lr_schedule.create() + _orig_log = wandb.log + + def _log(data, *args, step=None, **kwargs): + if step is not None and isinstance(data, dict) and 'loss' in data: + data = {**data, 'learning_rate': float(lr_fn(step))} + return _orig_log(data, *args, step=step, **kwargs) + + wandb.log = _log + openpi_train.main(cfg) + + +if __name__ == '__main__': + main() diff --git a/positronic/vendors/openpi/train.py b/positronic/vendors/openpi/train.py index 352e4897b..e5d106439 100644 --- a/positronic/vendors/openpi/train.py +++ b/positronic/vendors/openpi/train.py @@ -41,7 +41,8 @@ def main( local_input_dir = input_dir command = [uv_path, 'run', '--frozen', '--project', str(openpi_root), '--'] - command.extend(['python', 'scripts/train.py']) + launcher = Path(__file__).parent / '_launch.py' + command.extend(['python', str(launcher), '--openpi-root', openpi_root.as_posix()]) command.extend([config_name, '--exp-name', exp_name]) command.append('--resume' if resume else '--overwrite') if num_train_steps is not None: From bf5fa6a2fed6b0bdbd68beed525931a7181e8fa1 Mon Sep 17 00:00:00 2001 From: Sergey Arkhangelskiy Date: Mon, 18 May 2026 14:06:46 +0300 Subject: [PATCH 2/3] Patch `Run.log` so injected LR survives openpi's `wandb.init()` --- positronic/vendors/openpi/_launch.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/positronic/vendors/openpi/_launch.py b/positronic/vendors/openpi/_launch.py index 94276933e..883da96b6 100644 --- a/positronic/vendors/openpi/_launch.py +++ b/positronic/vendors/openpi/_launch.py @@ -3,8 +3,8 @@ import sys import openpi.training.config as _config -import wandb from openpi.training.optimizer import CosineDecaySchedule +from wandb.sdk.wandb_run import Run def _extract_openpi_root() -> str: @@ -35,14 +35,19 @@ def main(): ) lr_fn = cfg.lr_schedule.create() - _orig_log = wandb.log - def _log(data, *args, step=None, **kwargs): + # openpi calls `wandb.init()` inside its train main, which rebinds the + # module-level `wandb.log`. Patch the `Run.log` chokepoint instead so the + # injected learning_rate survives init regardless of patch timing. + _orig_log = Run.log + + def _log(self, data, *args, **kwargs): + step = kwargs.get('step', args[0] if args else None) if step is not None and isinstance(data, dict) and 'loss' in data: data = {**data, 'learning_rate': float(lr_fn(step))} - return _orig_log(data, *args, step=step, **kwargs) + return _orig_log(self, data, *args, **kwargs) - wandb.log = _log + Run.log = _log openpi_train.main(cfg) From bb84667e05d8976325fe0cf0dc537b90f5c75ba8 Mon Sep 17 00:00:00 2001 From: Sergey Arkhangelskiy Date: Mon, 18 May 2026 20:31:51 +0300 Subject: [PATCH 3/3] Drop `_launch.py`; pass `--lr-schedule.decay-steps` instead openpi PR #6 (merged) logs `learning_rate` natively, so the wandb monkeypatch is unnecessary. The decay-steps fix needs no wrapper either: `lr_schedule` is CLI-exposed, so pass `--lr-schedule.decay-steps` directly. --- positronic/vendors/openpi/_launch.py | 55 ---------------------------- positronic/vendors/openpi/train.py | 6 ++- 2 files changed, 4 insertions(+), 57 deletions(-) delete mode 100644 positronic/vendors/openpi/_launch.py diff --git a/positronic/vendors/openpi/_launch.py b/positronic/vendors/openpi/_launch.py deleted file mode 100644 index 883da96b6..000000000 --- a/positronic/vendors/openpi/_launch.py +++ /dev/null @@ -1,55 +0,0 @@ -import dataclasses -import importlib.util -import sys - -import openpi.training.config as _config -from openpi.training.optimizer import CosineDecaySchedule -from wandb.sdk.wandb_run import Run - - -def _extract_openpi_root() -> str: - argv = sys.argv[1:] - if len(argv) < 2 or argv[0] != '--openpi-root': - raise SystemExit('_launch.py expects `--openpi-root ` as its first argument') - root = argv[1] - sys.argv = [sys.argv[0], *argv[2:]] - return root - - -def _load_openpi_train(openpi_root: str): - spec = importlib.util.spec_from_file_location('openpi_train', f'{openpi_root}/scripts/train.py') - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -def main(): - openpi_root = _extract_openpi_root() - openpi_train = _load_openpi_train(openpi_root) - - cfg = _config.cli() - - if isinstance(cfg.lr_schedule, CosineDecaySchedule): - cfg = dataclasses.replace( - cfg, lr_schedule=dataclasses.replace(cfg.lr_schedule, decay_steps=cfg.num_train_steps) - ) - - lr_fn = cfg.lr_schedule.create() - - # openpi calls `wandb.init()` inside its train main, which rebinds the - # module-level `wandb.log`. Patch the `Run.log` chokepoint instead so the - # injected learning_rate survives init regardless of patch timing. - _orig_log = Run.log - - def _log(self, data, *args, **kwargs): - step = kwargs.get('step', args[0] if args else None) - if step is not None and isinstance(data, dict) and 'loss' in data: - data = {**data, 'learning_rate': float(lr_fn(step))} - return _orig_log(self, data, *args, **kwargs) - - Run.log = _log - openpi_train.main(cfg) - - -if __name__ == '__main__': - main() diff --git a/positronic/vendors/openpi/train.py b/positronic/vendors/openpi/train.py index e5d106439..ebf2b4727 100644 --- a/positronic/vendors/openpi/train.py +++ b/positronic/vendors/openpi/train.py @@ -41,12 +41,14 @@ def main( local_input_dir = input_dir command = [uv_path, 'run', '--frozen', '--project', str(openpi_root), '--'] - launcher = Path(__file__).parent / '_launch.py' - command.extend(['python', str(launcher), '--openpi-root', openpi_root.as_posix()]) + command.extend(['python', 'scripts/train.py']) command.extend([config_name, '--exp-name', exp_name]) command.append('--resume' if resume else '--overwrite') if num_train_steps is not None: command.append(f'--num-train-steps={num_train_steps}') + # Tie the cosine LR schedule to the run length so it spans the whole + # run instead of openpi's hardcoded default decay_steps. + command.append(f'--lr-schedule.decay-steps={num_train_steps}') command.extend(['--assets-base-dir', stats_dir.as_posix()]) command.extend(['--checkpoint-base-dir', output_dir.parent.parent.as_posix()]) command.extend(extra_args)