Fix repeated FLUX image generation memory lifecycle

This commit is contained in:
Mikei386
2026-09-20 22:12:55 +02:00
parent c78a083d4c
commit 3cbcf41fab
3 changed files with 137 additions and 14 deletions
+61
View File
@@ -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)