Fix repeated FLUX image generation memory lifecycle
This commit is contained in:
1 parent
c78a083d4c
commit
3cbcf41fab
3 files changed
+137
-14
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 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,19 +177,21 @@ 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)
|
||||||
transformer = Flux2Transformer2DModel.from_single_file(
|
with _install_fp8_converter() as scales:
|
||||||
TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer",
|
transformer = Flux2Transformer2DModel.from_single_file(
|
||||||
quantization_config=quantization, torch_dtype=torch.bfloat16,
|
TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer",
|
||||||
device_map={"": tx_index}, local_files_only=True)
|
quantization_config=quantization, torch_dtype=torch.bfloat16,
|
||||||
|
device_map={"": tx_index}, local_files_only=True)
|
||||||
patched = 0
|
patched = 0
|
||||||
for module_name, module in transformer.named_modules():
|
for module_name, module in transformer.named_modules():
|
||||||
if module_name not in scales:
|
if module_name not in scales:
|
||||||
@@ -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()
|
||||||
Reference in new issue
Block a user