Fix repeated FLUX image generation memory lifecycle
This commit is contained in:
@@ -0,0 +1,61 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user