| 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() |
|
|