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
+44 -14
View File
@@ -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()