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)
+32
View File
@@ -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/`.
+40 -10
View File
@@ -14,6 +14,8 @@ import json
import os import os
import signal import signal
import time import time
from contextlib import contextmanager
from threading import Lock
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path from pathlib import Path
from types import MethodType from types import MethodType
@@ -26,6 +28,7 @@ TRANSFORMER_FILE = os.environ.get(
"FLUX_TRANSFORMER_FILE", "/models/fp8/flux-2-klein-9b-fp8.safetensors") "FLUX_TRANSFORMER_FILE", "/models/fp8/flux-2-klein-9b-fp8.safetensors")
OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve() OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve()
ACTIVE = False ACTIVE = False
GENERATION_LOCK = Lock()
os.environ.setdefault("DIFFUSERS_VERBOSITY", "error") os.environ.setdefault("DIFFUSERS_VERBOSITY", "error")
os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error") os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error")
@@ -50,6 +53,7 @@ def _devices(torch):
torch.device(f"cuda:{encoder_index}")) torch.device(f"cuda:{encoder_index}"))
@contextmanager
def _install_fp8_converter(): def _install_fp8_converter():
import diffusers.loaders.single_file_model as single_file_model 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[ single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = convert "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): def _fp8_forward(torch, module, inputs):
@@ -123,6 +131,19 @@ def _fp8_forward(torch, module, inputs):
def generate(data: dict) -> dict: def generate(data: dict) -> dict:
global ACTIVE global ACTIVE
import torch 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, from diffusers import (Flux2KleinPipeline, Flux2Transformer2DModel,
NVIDIAModelOptConfig) NVIDIAModelOptConfig)
from modelopt.torch.opt import enable_huggingface_checkpointing from modelopt.torch.opt import enable_huggingface_checkpointing
@@ -156,15 +177,17 @@ def generate(data: dict) -> dict:
source_images.append(opened.convert("RGB")) source_images.append(opened.convert("RGB"))
started = time.monotonic() started = time.monotonic()
ACTIVE = True
transformer = text_encoder = pipe = latent = decoded = image = None transformer = text_encoder = pipe = latent = decoded = image = None
prompt_embeds = generator = module = None
kwargs = {}
scales = {}
try: try:
enable_huggingface_checkpointing() enable_huggingface_checkpointing()
scales = _install_fp8_converter()
tx_index, enc_index, tx_device, enc_device = _devices(torch) tx_index, enc_index, tx_device, enc_device = _devices(torch)
quantization = NVIDIAModelOptConfig( quantization = NVIDIAModelOptConfig(
quant_type="FP8", weight_only=False, quant_type="FP8", weight_only=False,
modelopt_config=FP8_DEFAULT_CFG) modelopt_config=FP8_DEFAULT_CFG)
with _install_fp8_converter() as scales:
transformer = Flux2Transformer2DModel.from_single_file( transformer = Flux2Transformer2DModel.from_single_file(
TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer", TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer",
quantization_config=quantization, torch_dtype=torch.bfloat16, quantization_config=quantization, torch_dtype=torch.bfloat16,
@@ -182,6 +205,8 @@ def generate(data: dict) -> dict:
patched += 1 patched += 1
if patched != len(scales): if patched != len(scales):
raise RuntimeError(f"patched only {patched} of {len(scales)} FP8 layers") raise RuntimeError(f"patched only {patched} of {len(scales)} FP8 layers")
module = None
scales.clear()
transformer.to(tx_device) transformer.to(tx_device)
text_encoder = Qwen3ForCausalLM.from_pretrained( text_encoder = Qwen3ForCausalLM.from_pretrained(
@@ -219,8 +244,8 @@ def generate(data: dict) -> dict:
latent = pipe(**kwargs).images latent = pipe(**kwargs).images
pipe.transformer = None pipe.transformer = None
del transformer, text_encoder, prompt_embeds, generator kwargs.clear()
transformer = text_encoder = None transformer = text_encoder = prompt_embeds = generator = None
gc.collect() gc.collect()
torch.cuda.empty_cache() torch.cuda.empty_cache()
latent = latent.to(device=tx_device, dtype=pipe.vae.dtype) latent = latent.to(device=tx_device, dtype=pipe.vae.dtype)
@@ -234,12 +259,16 @@ def generate(data: dict) -> dict:
"load_seconds": round(loaded, 3), "load_seconds": round(loaded, 3),
"model": "FLUX.2-klein-9B-fp8-beta"} "model": "FLUX.2-klein-9B-fp8-beta"}
finally: finally:
for value in (image, decoded, latent, pipe, text_encoder, transformer): # Drop the actual references before collecting bound-method cycles.
if value is not None: kwargs.clear()
del value 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() gc.collect()
torch.cuda.empty_cache() torch.cuda.empty_cache()
ACTIVE = False
class Handler(BaseHTTPRequestHandler): class Handler(BaseHTTPRequestHandler):
@@ -279,4 +308,5 @@ class Handler(BaseHTTPRequestHandler):
self.reply(500, {"status": "error", "message": str(exc)}) self.reply(500, {"status": "error", "message": str(exc)})
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever() if __name__ == "__main__":
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()