Files
AI-Profile-Router/dev/test_image_worker_lifecycle.py

62 lines
3.0 KiB
Python

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