62 lines
3.0 KiB
Python
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)
|