Download scripts/test_wandb_fallback.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 2.62 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_wandb_fallback.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/test_wandb_fallback.py
-
curl -L -o test_wandb_fallback.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/test_wandb_fallback.py
2.62 kB
| """CPU-only regression tests for W&B fallback with Cosmos's frozen config.""" | |
| import unittest | |
| from unittest.mock import Mock, patch | |
| import wandb | |
| from cosmos_framework.utils.config import Config, JobConfig | |
| from scripts.train_widowx_zeva import with_wandb_offline_fallback | |
| def frozen_config(mode='online'): | |
| config = Config( | |
| model={}, optimizer={}, scheduler={}, dataloader_train=None, | |
| dataloader_val=None, | |
| job=JobConfig(project='Zeva', group='widowx', name='ipec_personal', | |
| wandb_mode=mode), | |
| ) | |
| config.freeze() | |
| return config | |
| class WandbFallbackTest(unittest.TestCase): | |
| def test_online_success_does_not_retry(self): | |
| initialize = Mock(return_value='online') | |
| config = frozen_config() | |
| self.assertEqual(with_wandb_offline_fallback(initialize)(config, None), 'online') | |
| initialize.assert_called_once_with(config, None) | |
| def test_retry_preserves_identity_and_frozen_config(self): | |
| initialize = Mock(side_effect=[wandb.errors.CommError('network failure'), 'offline']) | |
| config = frozen_config() | |
| with patch('wandb.teardown') as teardown: | |
| self.assertEqual(with_wandb_offline_fallback(initialize)(config, None), 'offline') | |
| teardown.assert_called_once_with() | |
| retry = initialize.call_args.args[0] | |
| self.assertEqual(retry.job.wandb_mode, 'offline') | |
| self.assertEqual(retry.job.path_local, config.job.path_local) | |
| self.assertIs(retry.checkpoint, config.checkpoint) | |
| self.assertEqual(config.job.wandb_mode, 'online') | |
| self.assertTrue(config._is_frozen and config.job._is_frozen) | |
| def test_offline_failure_is_not_retried(self): | |
| initialize = Mock(side_effect=wandb.errors.CommError('offline failure')) | |
| with self.assertRaises(wandb.errors.CommError): | |
| with_wandb_offline_fallback(initialize)(frozen_config('offline'), None) | |
| self.assertEqual(initialize.call_count, 1) | |
| def test_unrelated_error_is_not_hidden(self): | |
| initialize = Mock(side_effect=RuntimeError('training failure')) | |
| with self.assertRaises(RuntimeError): | |
| with_wandb_offline_fallback(initialize)(frozen_config(), None) | |
| self.assertEqual(initialize.call_count, 1) | |
| def test_failed_retry_propagates(self): | |
| initialize = Mock(side_effect=wandb.errors.CommError('unavailable')) | |
| with self.assertRaises(wandb.errors.CommError): | |
| with_wandb_offline_fallback(initialize)(frozen_config(), None) | |
| self.assertEqual(initialize.call_count, 2) | |
| if __name__ == '__main__': | |
| unittest.main() | |