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