import importlib import unittest import os from pathlib import Path from unittest.mock import patch from omegaconf import OmegaConf from experiments import exp_base class DummyExperiment(exp_base.BaseLightningExperiment): compatible_algorithms = {} compatible_datasets = {} def _build_training_loader(self): return None def _build_validation_loader(self): return None def _build_test_loader(self): return None class DummyTrainer: init_calls = [] fit_calls = [] validate_calls = [] test_calls = [] save_checkpoint_calls = [] def __init__(self, *args, **kwargs): self.init_calls.append({"args": args, "kwargs": kwargs}) def fit(self, *args, **kwargs): self.fit_calls.append({"args": args, "kwargs": kwargs}) def validate(self, *args, **kwargs): self.validate_calls.append({"args": args, "kwargs": kwargs}) def test(self, *args, **kwargs): self.test_calls.append({"args": args, "kwargs": kwargs}) def save_checkpoint(self, path): self.save_checkpoint_calls.append(path) def _root_cfg(auto_resuming=True, algorithm_name="dummy", zero_init_gate=False, only_tune_memory=False): return OmegaConf.create( { "debug": False, "_auto_resuming": auto_resuming, "customized_load": True, "seperate_load": True, "zero_init_gate": zero_init_gate, "only_tune_memory": only_tune_memory, "diffusion_model_path": "oasis500m.safetensors", "vae_path": "vit-l-20.safetensors", "algorithm": {"_name": algorithm_name}, "dataset": {"_name": "dummy"}, "experiment": { "debug": False, "num_nodes": 1, "tasks": ["training"], "training": { "compile": False, "precision": 32, "batch_size": 1, "max_epochs": 1, "max_steps": 1, "max_time": None, "data": {"shuffle": False, "num_workers": 0}, "optim": {"gradient_clip_val": 0, "accumulate_grad_batches": 1}, }, "validation": { "compile": False, "precision": 32, "inference_mode": True, "val_every_n_step": None, "val_every_n_epoch": None, "limit_batch": 0, "batch_size": 1, "data": {"shuffle": False, "num_workers": 0}, }, "test": { "compile": False, "precision": 32, "inference_mode": True, "limit_batch": 0, "batch_size": 1, "data": {"shuffle": False, "num_workers": 0}, }, }, } ) class ResumeCheckpointLogicTests(unittest.TestCase): def test_dememwm_rejects_zero_init_gate(self): with self.assertRaisesRegex(ValueError, "zero_init_gate.*dememwm_base"): DummyExperiment(_root_cfg(algorithm_name="dememwm_base", zero_init_gate=True), logger=None) def test_dememwm_rejects_only_tune_memory(self): with self.assertRaisesRegex(ValueError, "only_tune_memory.*dememwm_base"): DummyExperiment(_root_cfg(algorithm_name="dememwm_base", only_tune_memory=True), logger=None) def test_non_dememwm_keeps_stale_flag_behavior(self): experiment = DummyExperiment( _root_cfg(algorithm_name="dummy", zero_init_gate=True, only_tune_memory=True), logger=None, ) self.assertTrue(experiment.zero_init_gate) self.assertTrue(experiment.only_tune_memory) def test_auto_resume_takes_priority_over_custom_separate_load(self): DummyTrainer.fit_calls = [] experiment = DummyExperiment(_root_cfg(auto_resuming=True), logger=None, ckpt_path="/tmp/last.ckpt") experiment.algo = object() with patch.object(exp_base.pl, "Trainer", DummyTrainer), patch.object( exp_base, "load_custom_checkpoint", side_effect=AssertionError("custom load should not run during auto-resume"), ): experiment.training() self.assertEqual(DummyTrainer.fit_calls[0]["kwargs"]["ckpt_path"], "/tmp/last.ckpt") def test_curriculum_updates_stage_config_and_checkpoint_handoff(self): DummyTrainer.init_calls = [] DummyTrainer.fit_calls = [] DummyTrainer.save_checkpoint_calls = [] cfg = _root_cfg(auto_resuming=True) cfg.experiment.training.curriculum = { "enabled": True, "stages": [ { "name": "near", "until_step": 3, "dataset": {"wo_updown": True, "memory_selection": {"pose_similarity_radius": 2.0}}, "algorithm": {"memory_selection": {"pose_similarity_radius": 2.0}}, }, { "name": "full", "until_step": 5, "dataset": {"wo_updown": False, "memory_selection": {"pose_similarity_radius": 8.0}}, "algorithm": {"memory_selection": {"pose_similarity_radius": 8.0}}, }, ], } experiment = DummyExperiment(cfg, logger=None, ckpt_path="/tmp/last.ckpt") experiment.algo = object() seen_dataset = [] def build_training_loader(): seen_dataset.append( ( bool(experiment.root_cfg.dataset.wo_updown), float(experiment.root_cfg.dataset.memory_selection.pose_similarity_radius), ) ) return None experiment._build_training_loader = build_training_loader with patch.object(exp_base.pl, "Trainer", DummyTrainer), patch.object( DummyExperiment, "_curriculum_checkpoint_path", return_value=Path("/tmp/curriculum_stage.ckpt"), ): experiment.training() self.assertEqual(seen_dataset, [(True, 2.0), (False, 8.0)]) self.assertEqual( [call["kwargs"]["ckpt_path"] for call in DummyTrainer.fit_calls], ["/tmp/last.ckpt", Path("/tmp/curriculum_stage.ckpt")], ) self.assertEqual( [call["kwargs"]["max_steps"] for call in DummyTrainer.init_calls], [3, 5], ) self.assertEqual(DummyTrainer.save_checkpoint_calls, [Path("/tmp/curriculum_stage.ckpt")]) self.assertEqual(float(experiment.root_cfg.algorithm.memory_selection.pose_similarity_radius), 8.0) def test_auto_resume_curriculum_skips_completed_stages(self): DummyTrainer.init_calls = [] DummyTrainer.fit_calls = [] DummyTrainer.save_checkpoint_calls = [] cfg = _root_cfg(auto_resuming=True) cfg.experiment.training.curriculum = { "enabled": True, "stages": [ {"name": "near", "until_step": 3, "dataset": {"wo_updown": True}}, {"name": "full", "until_step": 5, "dataset": {"wo_updown": False}}, ], } experiment = DummyExperiment(cfg, logger=None, ckpt_path="/tmp/epoch1_step4.ckpt") experiment.algo = object() seen_dataset = [] def build_training_loader(): seen_dataset.append(bool(experiment.root_cfg.dataset.wo_updown)) return None experiment._build_training_loader = build_training_loader with patch.object(exp_base.pl, "Trainer", DummyTrainer), patch.object( DummyExperiment, "_curriculum_checkpoint_path", return_value=Path("/tmp/curriculum_stage.ckpt"), ): experiment.training() self.assertEqual(seen_dataset, [False]) self.assertEqual( [call["kwargs"]["ckpt_path"] for call in DummyTrainer.fit_calls], ["/tmp/epoch1_step4.ckpt"], ) self.assertEqual( [call["kwargs"]["max_steps"] for call in DummyTrainer.init_calls], [5], ) self.assertEqual(DummyTrainer.save_checkpoint_calls, []) def test_validation_load_checkpoint_takes_priority_over_separate_load(self): DummyTrainer.validate_calls = [] experiment = DummyExperiment(_root_cfg(auto_resuming=False), logger=None, ckpt_path="/tmp/trained.ckpt") experiment.algo = object() with patch.object(exp_base.pl, "Trainer", DummyTrainer), patch.object(exp_base, "load_custom_checkpoint") as load_mock: experiment.validation() load_mock.assert_called_once_with(algo=experiment.algo, checkpoint_path="/tmp/trained.ckpt") self.assertIsNone(DummyTrainer.validate_calls[0]["kwargs"]["ckpt_path"]) def test_test_load_checkpoint_takes_priority_over_separate_load(self): DummyTrainer.test_calls = [] experiment = DummyExperiment(_root_cfg(auto_resuming=False), logger=None, ckpt_path="/tmp/trained.ckpt") experiment.algo = object() with patch.object(exp_base.pl, "Trainer", DummyTrainer), patch.object(exp_base, "load_custom_checkpoint") as load_mock: experiment.test() load_mock.assert_called_once_with(algo=experiment.algo, checkpoint_path="/tmp/trained.ckpt") self.assertIsNone(DummyTrainer.test_calls[0]["kwargs"]["ckpt_path"]) def test_pre_lightning_rank_uses_slurm_procid(self): import main import utils.distributed_utils as distributed_utils old_env = os.environ.copy() try: os.environ.pop("RANK", None) os.environ["SLURM_PROCID"] = "1" self.assertEqual(main._process_rank(), 1) reloaded = importlib.reload(distributed_utils) self.assertFalse(reloaded.is_rank_zero) finally: os.environ.clear() os.environ.update(old_env) importlib.reload(distributed_utils) if __name__ == "__main__": unittest.main()