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
@@ -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()
|
||||
Reference in new issue
Block a user