"""GPU-free regressions for sequential image requests.""" import importlib.util import os from pathlib import Path import sys import types import unittest from unittest.mock import Mock, patch from contextlib import nullcontext class ImageWorkerLifecycleTests(unittest.TestCase): @classmethod def setUpClass(cls): path = Path(__file__).resolve().parents[1] / 'platform/docker/image-worker/image_worker_9b.py' spec = importlib.util.spec_from_file_location('image_worker_lifecycle', path) cls.worker = importlib.util.module_from_spec(spec) with patch.dict(os.environ, {'WORKER_TOKEN': 'test-' * 10}), patch('signal.signal'): spec.loader.exec_module(cls.worker) def test_converter_is_restored_and_second_request_keeps_scales(self): original = Mock(return_value={}) entry = {'checkpoint_mapping_fn': original} sfm = types.ModuleType('diffusers.loaders.single_file_model') sfm.SINGLE_FILE_LOADABLE_CLASSES = {'Flux2Transformer2DModel': entry} diffusers = types.ModuleType('diffusers') loaders = types.ModuleType('diffusers.loaders') diffusers.loaders = loaders loaders.single_file_model = sfm with patch.dict(sys.modules, {'diffusers': diffusers, 'diffusers.loaders': loaders, 'diffusers.loaders.single_file_model': sfm}): for _ in range(2): with self.worker._install_fp8_converter() as scales: entry['checkpoint_mapping_fn']({'double_blocks.0.img_attn.proj.input_scale': Mock()}) self.assertIn('transformer_blocks.0.attn.to_out.0', scales) self.assertIs(entry['checkpoint_mapping_fn'], original) with self.assertRaises(RuntimeError): with self.worker._install_fp8_converter(): raise RuntimeError('load failed') self.assertIs(entry['checkpoint_mapping_fn'], original) def test_inference_context_and_cleanup_on_success_and_failure(self): torch = Mock() torch.inference_mode.side_effect = lambda: nullcontext() for fails in (False, True): def run(data, module): self.assertTrue(self.worker.ACTIVE) self.assertTrue(self.worker.GENERATION_LOCK.locked()) if fails: raise RuntimeError('CUDA OOM') return {'status': 'ok'} with patch.dict(sys.modules, {'torch': torch}), patch.object(self.worker, '_generate', side_effect=run): if fails: with self.assertRaises(RuntimeError): self.worker.generate({}) else: self.assertEqual(self.worker.generate({}), {'status': 'ok'}) self.assertFalse(self.worker.ACTIVE) self.assertFalse(self.worker.GENERATION_LOCK.locked()) self.assertEqual(torch.inference_mode.call_count, 2) self.assertEqual(torch.cuda.empty_cache.call_count, 2)