From 3cbcf41fabe7ed864f9277a6ac80a4ebf7524604 Mon Sep 17 00:00:00 2001 From: Mikei386 <44135113+Mikei386@users.noreply.github.com> Date: Sun, 20 Sep 2026 22:12:55 +0200 Subject: [PATCH] Fix repeated FLUX image generation memory lifecycle --- dev/test_image_worker_lifecycle.py | 61 +++++++++++++++++++ docs/FLUX_REPEAT_FIX_20260920.md | 32 ++++++++++ .../docker/image-worker/image_worker_9b.py | 58 +++++++++++++----- 3 files changed, 137 insertions(+), 14 deletions(-) create mode 100644 dev/test_image_worker_lifecycle.py create mode 100644 docs/FLUX_REPEAT_FIX_20260920.md diff --git a/dev/test_image_worker_lifecycle.py b/dev/test_image_worker_lifecycle.py new file mode 100644 index 0000000..85d93b8 --- /dev/null +++ b/dev/test_image_worker_lifecycle.py @@ -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) diff --git a/docs/FLUX_REPEAT_FIX_20260920.md b/docs/FLUX_REPEAT_FIX_20260920.md new file mode 100644 index 0000000..96fe9b0 --- /dev/null +++ b/docs/FLUX_REPEAT_FIX_20260920.md @@ -0,0 +1,32 @@ +# FLUX-Wiederholungsfehler am 20.09.2026 + +Beim OpenClaw-Bildauftrag erzeugte der Worker ein Bild erfolgreich und scheiterte +beim folgenden Bild mit CUDA OOM (angeforderte 4,23 GiB bei 3,73 GiB frei). +Der Router hatte Qwen und TTS korrekt gestoppt. Kein Benchmarkcontainer war aktiv. +Die experimentelle Qwen-Mindestreserve war nicht die Ursache. + +## Korrektur + +- Den globalen FP8-Checkpoint-Konverter nur während des Ladens ersetzen und auch + bei Fehlern wiederherstellen. Die bisher verschachtelten Wrapper entfernten + beim zweiten Laden Skalierungswerte, bevor der aktuelle Wrapper sie verwenden + konnte. Ein Regressionstest deckt zwei Ladevorgänge und den Fehlerpfad ab. +- Den gesamten Tensorpfad einschließlich direkter Textencoder- und VAE-Aufrufe + unter `torch.inference_mode()` ausführen. +- Echte Modell-/Tensorreferenzen und Argument-Dictionaries vor GC und + `empty_cache()` freigeben. `del value` auf einer Schleifenvariablen hatte die + eigentlichen lokalen Referenzen nicht entfernt. +- Generierungen im Worker serialisieren; Health-Endpunkt bleibt erreichbar. + +## Validierung und Betrieb + +88 GPU-freie Tests bestanden. Ein echter Router-Auftrag mit `n=2`, 1024×1024, +vier Schritten und Seed 12345 erzeugte beide Bilder im selben Worker-Prozess; +beide `/generate`-Aufrufe lieferten HTTP 200. Die Modellgewichte und +Quantisierung bleiben unverändert. Der Test prüft Wiederholbarkeit des Ablaufs, +nicht die Gleichheit mit Bildern vor der Reparatur. + +Nur das Worker-Image wurde aktualisiert, auf Basis der vorhandenen Abhängigkeiten +(kein Paket-, Treiber- oder Kernelupdate). Das vorherige Image bleibt als +`mike-ai/image-worker:before-repeat-fix-20260920` für Rollback vorhanden. +Remote-Testunterlagen: `/data/benchmarks/flux-repeat-fix-20260920/`. diff --git a/platform/docker/image-worker/image_worker_9b.py b/platform/docker/image-worker/image_worker_9b.py index d86cfdd..c41ae9d 100644 --- a/platform/docker/image-worker/image_worker_9b.py +++ b/platform/docker/image-worker/image_worker_9b.py @@ -14,6 +14,8 @@ import json import os import signal import time +from contextlib import contextmanager +from threading import Lock from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from types import MethodType @@ -26,6 +28,7 @@ TRANSFORMER_FILE = os.environ.get( "FLUX_TRANSFORMER_FILE", "/models/fp8/flux-2-klein-9b-fp8.safetensors") OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve() ACTIVE = False +GENERATION_LOCK = Lock() os.environ.setdefault("DIFFUSERS_VERBOSITY", "error") os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error") @@ -50,6 +53,7 @@ def _devices(torch): torch.device(f"cuda:{encoder_index}")) +@contextmanager def _install_fp8_converter(): import diffusers.loaders.single_file_model as single_file_model @@ -99,7 +103,11 @@ def _install_fp8_converter(): single_file_model.SINGLE_FILE_LOADABLE_CLASSES[ "Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = convert - return scales + try: + yield scales + finally: + single_file_model.SINGLE_FILE_LOADABLE_CLASSES[ + "Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = original def _fp8_forward(torch, module, inputs): @@ -123,6 +131,19 @@ def _fp8_forward(torch, module, inputs): def generate(data: dict) -> dict: global ACTIVE import torch + + with GENERATION_LOCK: + ACTIVE = True + try: + with torch.inference_mode(): + return _generate(data, torch) + finally: + gc.collect() + torch.cuda.empty_cache() + ACTIVE = False + + +def _generate(data: dict, torch) -> dict: from diffusers import (Flux2KleinPipeline, Flux2Transformer2DModel, NVIDIAModelOptConfig) from modelopt.torch.opt import enable_huggingface_checkpointing @@ -156,19 +177,21 @@ def generate(data: dict) -> dict: source_images.append(opened.convert("RGB")) started = time.monotonic() - ACTIVE = True transformer = text_encoder = pipe = latent = decoded = image = None + prompt_embeds = generator = module = None + kwargs = {} + scales = {} try: enable_huggingface_checkpointing() - scales = _install_fp8_converter() tx_index, enc_index, tx_device, enc_device = _devices(torch) quantization = NVIDIAModelOptConfig( quant_type="FP8", weight_only=False, modelopt_config=FP8_DEFAULT_CFG) - transformer = Flux2Transformer2DModel.from_single_file( - TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer", - quantization_config=quantization, torch_dtype=torch.bfloat16, - device_map={"": tx_index}, local_files_only=True) + with _install_fp8_converter() as scales: + transformer = Flux2Transformer2DModel.from_single_file( + TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer", + quantization_config=quantization, torch_dtype=torch.bfloat16, + device_map={"": tx_index}, local_files_only=True) patched = 0 for module_name, module in transformer.named_modules(): if module_name not in scales: @@ -182,6 +205,8 @@ def generate(data: dict) -> dict: patched += 1 if patched != len(scales): raise RuntimeError(f"patched only {patched} of {len(scales)} FP8 layers") + module = None + scales.clear() transformer.to(tx_device) text_encoder = Qwen3ForCausalLM.from_pretrained( @@ -219,8 +244,8 @@ def generate(data: dict) -> dict: latent = pipe(**kwargs).images pipe.transformer = None - del transformer, text_encoder, prompt_embeds, generator - transformer = text_encoder = None + kwargs.clear() + transformer = text_encoder = prompt_embeds = generator = None gc.collect() torch.cuda.empty_cache() latent = latent.to(device=tx_device, dtype=pipe.vae.dtype) @@ -234,12 +259,16 @@ def generate(data: dict) -> dict: "load_seconds": round(loaded, 3), "model": "FLUX.2-klein-9B-fp8-beta"} finally: - for value in (image, decoded, latent, pipe, text_encoder, transformer): - if value is not None: - del value + # Drop the actual references before collecting bound-method cycles. + kwargs.clear() + scales.clear() + image = decoded = latent = pipe = text_encoder = transformer = None + prompt_embeds = generator = module = None + for source_image in source_images: + source_image.close() + source_images.clear() gc.collect() torch.cuda.empty_cache() - ACTIVE = False class Handler(BaseHTTPRequestHandler): @@ -279,4 +308,5 @@ class Handler(BaseHTTPRequestHandler): self.reply(500, {"status": "error", "message": str(exc)}) -ThreadingHTTPServer((HOST, PORT), Handler).serve_forever() +if __name__ == "__main__": + ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()