Fix repeated FLUX image generation memory lifecycle
This commit is contained in:
1 parent
c78a083d4c
commit
3cbcf41fab
3 files changed
+133
-10
No files matched your search
@@ -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)
|
||||
@@ -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/`.
|
||||
@@ -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,15 +177,17 @@ 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)
|
||||
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,
|
||||
@@ -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()
|
||||
Reference in new issue
Block a user