DynaTTT / scripts /test_wandb_fallback.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
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()