Synchronize repository with Athena deployment
This commit is contained in:
@@ -0,0 +1,282 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Private FLUX.2 Klein 9B FP8 beta worker for Athena's two GPUs.
|
||||
|
||||
The FP8 diffusion transformer runs on the RTX 5080. A Qwen3-8B NF4 text
|
||||
encoder runs on the RTX 3060 while the profile controller temporarily pauses
|
||||
Qwen3-TTS. The transformer and encoder are released before VAE decoding so
|
||||
the 1024px decoder has sufficient workspace on the RTX 5080.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import json
|
||||
import os
|
||||
import signal
|
||||
import time
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from types import MethodType
|
||||
|
||||
HOST = os.environ.get("WORKER_HOST", "0.0.0.0")
|
||||
PORT = int(os.environ.get("WORKER_PORT", "8086"))
|
||||
TOKEN = os.environ.get("WORKER_TOKEN", "").strip()
|
||||
COMPONENT_DIR = os.environ.get("FLUX_COMPONENT_DIR", "/models/components")
|
||||
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
|
||||
|
||||
os.environ.setdefault("DIFFUSERS_VERBOSITY", "error")
|
||||
os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error")
|
||||
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
|
||||
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||||
|
||||
if len(TOKEN) < 32:
|
||||
raise RuntimeError("WORKER_TOKEN is missing or too short")
|
||||
|
||||
signal.signal(signal.SIGTERM, lambda *_: os._exit(0))
|
||||
|
||||
|
||||
def _devices(torch):
|
||||
if torch.cuda.device_count() != 2:
|
||||
raise RuntimeError("FLUX 9B beta requires exactly two visible CUDA GPUs")
|
||||
totals = {i: torch.cuda.get_device_properties(i).total_memory
|
||||
for i in range(torch.cuda.device_count())}
|
||||
transformer_index = max(totals, key=totals.get)
|
||||
encoder_index = min(totals, key=totals.get)
|
||||
return (transformer_index, encoder_index,
|
||||
torch.device(f"cuda:{transformer_index}"),
|
||||
torch.device(f"cuda:{encoder_index}"))
|
||||
|
||||
|
||||
def _install_fp8_converter():
|
||||
import diffusers.loaders.single_file_model as single_file_model
|
||||
|
||||
original = single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
|
||||
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"]
|
||||
scales = {}
|
||||
double_map = {
|
||||
"img_attn.proj": "attn.to_out.0",
|
||||
"img_mlp.0": "ff.linear_in",
|
||||
"img_mlp.2": "ff.linear_out",
|
||||
"txt_attn.proj": "attn.to_add_out",
|
||||
"txt_mlp.0": "ff_context.linear_in",
|
||||
"txt_mlp.2": "ff_context.linear_out",
|
||||
}
|
||||
single_map = {
|
||||
"linear1": "attn.to_qkv_mlp_proj",
|
||||
"linear2": "attn.to_out",
|
||||
}
|
||||
|
||||
def record(key, value):
|
||||
parts = key.split(".")
|
||||
scale_name, block = parts[-1], parts[1]
|
||||
within = ".".join(parts[2:-1])
|
||||
if parts[0] == "double_blocks":
|
||||
if within == "img_attn.qkv":
|
||||
targets = ("attn.to_q", "attn.to_k", "attn.to_v")
|
||||
elif within == "txt_attn.qkv":
|
||||
targets = ("attn.add_q_proj", "attn.add_k_proj",
|
||||
"attn.add_v_proj")
|
||||
else:
|
||||
targets = (double_map[within],)
|
||||
prefix = f"transformer_blocks.{block}"
|
||||
elif parts[0] == "single_blocks":
|
||||
targets = (single_map[within],)
|
||||
prefix = f"single_transformer_blocks.{block}"
|
||||
else:
|
||||
raise ValueError(f"unexpected FP8 scale key: {key}")
|
||||
for target in targets:
|
||||
scales.setdefault(f"{prefix}.{target}", {})[scale_name] = value.clone()
|
||||
|
||||
def convert(checkpoint, **kwargs):
|
||||
scales.clear()
|
||||
for key in list(checkpoint):
|
||||
if key.endswith((".input_scale", ".weight_scale")):
|
||||
record(key, checkpoint.pop(key))
|
||||
return original(checkpoint=checkpoint, **kwargs)
|
||||
|
||||
single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
|
||||
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = convert
|
||||
return scales
|
||||
|
||||
|
||||
def _fp8_forward(torch, module, inputs):
|
||||
shape = inputs.shape
|
||||
input_fp8 = ((inputs / module._fp8_input_scale)
|
||||
.clamp(torch.finfo(torch.float8_e4m3fn).min,
|
||||
torch.finfo(torch.float8_e4m3fn).max)
|
||||
.to(torch.float8_e4m3fn).reshape(-1, shape[-1]))
|
||||
output = torch._scaled_mm(
|
||||
input_fp8,
|
||||
module.weight.reshape(-1, module.weight.shape[-1]).t(),
|
||||
scale_a=module._fp8_input_scale,
|
||||
scale_b=module._fp8_weight_scale,
|
||||
bias=module.bias,
|
||||
out_dtype=inputs.dtype,
|
||||
use_fast_accum=True,
|
||||
)
|
||||
return output.reshape(*shape[:-1], output.shape[-1])
|
||||
|
||||
|
||||
def generate(data: dict) -> dict:
|
||||
global ACTIVE
|
||||
import torch
|
||||
from diffusers import (Flux2KleinPipeline, Flux2Transformer2DModel,
|
||||
NVIDIAModelOptConfig)
|
||||
from modelopt.torch.opt import enable_huggingface_checkpointing
|
||||
from modelopt.torch.quantization.config import FP8_DEFAULT_CFG
|
||||
from PIL import Image
|
||||
from transformers import BitsAndBytesConfig, Qwen3ForCausalLM
|
||||
|
||||
prompt, filename = data.get("prompt"), data.get("filename")
|
||||
if not isinstance(prompt, str) or not prompt.strip() or len(prompt) > 8000:
|
||||
raise ValueError("invalid prompt")
|
||||
if (not isinstance(filename, str) or Path(filename).name != filename
|
||||
or not filename.endswith(".png")):
|
||||
raise ValueError("invalid filename")
|
||||
width, height = int(data.get("width", 1024)), int(data.get("height", 1024))
|
||||
if (width, height) != (1024, 1024):
|
||||
raise ValueError("FLUX 9B beta currently supports only 1024x1024")
|
||||
if int(data.get("steps", 4)) != 4 or float(data.get("guidance", 1.0)) != 1.0:
|
||||
raise ValueError("FLUX 9B beta requires steps=4 and guidance=1.0")
|
||||
|
||||
source_files = data.get("source_files") or []
|
||||
if not isinstance(source_files, list) or len(source_files) > 4:
|
||||
raise ValueError("invalid source image list")
|
||||
source_images = []
|
||||
for name in source_files:
|
||||
if not isinstance(name, str) or Path(name).name != name:
|
||||
raise ValueError("invalid source image filename")
|
||||
source = (OUTPUT_DIR / name).resolve()
|
||||
if source.parent != OUTPUT_DIR or not source.is_file():
|
||||
raise ValueError("source image not found")
|
||||
with Image.open(source) as opened:
|
||||
source_images.append(opened.convert("RGB"))
|
||||
|
||||
started = time.monotonic()
|
||||
ACTIVE = True
|
||||
transformer = text_encoder = pipe = latent = decoded = image = None
|
||||
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)
|
||||
patched = 0
|
||||
for module_name, module in transformer.named_modules():
|
||||
if module_name not in scales:
|
||||
continue
|
||||
module.register_buffer("_fp8_input_scale",
|
||||
scales[module_name]["input_scale"])
|
||||
module.register_buffer("_fp8_weight_scale",
|
||||
scales[module_name]["weight_scale"])
|
||||
module.forward = MethodType(
|
||||
lambda self, inputs: _fp8_forward(torch, self, inputs), module)
|
||||
patched += 1
|
||||
if patched != len(scales):
|
||||
raise RuntimeError(f"patched only {patched} of {len(scales)} FP8 layers")
|
||||
transformer.to(tx_device)
|
||||
|
||||
text_encoder = Qwen3ForCausalLM.from_pretrained(
|
||||
os.path.join(COMPONENT_DIR, "text_encoder"),
|
||||
torch_dtype=torch.bfloat16, low_cpu_mem_usage=True,
|
||||
quantization_config=BitsAndBytesConfig(
|
||||
load_in_4bit=True, bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
bnb_4bit_use_double_quant=True),
|
||||
device_map={"": enc_index}, local_files_only=True)
|
||||
pipe = Flux2KleinPipeline.from_pretrained(
|
||||
COMPONENT_DIR, transformer=transformer, text_encoder=text_encoder,
|
||||
torch_dtype=torch.bfloat16, local_files_only=True)
|
||||
pipe.vae.enable_slicing()
|
||||
pipe.vae.enable_tiling()
|
||||
pipe.vae.to(tx_device)
|
||||
loaded = time.monotonic() - started
|
||||
|
||||
prompt_embeds, _ = pipe.encode_prompt(
|
||||
prompt.strip(), device=enc_device, max_sequence_length=128)
|
||||
prompt_embeds = prompt_embeds.to(tx_device)
|
||||
pipe.text_encoder = None
|
||||
seed = data.get("seed")
|
||||
generator = None if seed is None else torch.Generator(
|
||||
device=tx_device).manual_seed(int(seed))
|
||||
kwargs = {
|
||||
"prompt": None, "prompt_embeds": prompt_embeds,
|
||||
"height": height, "width": width, "num_inference_steps": 4,
|
||||
"guidance_scale": 1.0, "generator": generator,
|
||||
"output_type": "latent",
|
||||
}
|
||||
if source_images:
|
||||
kwargs["image"] = (source_images[0] if len(source_images) == 1
|
||||
else source_images)
|
||||
latent = pipe(**kwargs).images
|
||||
|
||||
pipe.transformer = None
|
||||
del transformer, text_encoder, prompt_embeds, generator
|
||||
transformer = text_encoder = None
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
latent = latent.to(device=tx_device, dtype=pipe.vae.dtype)
|
||||
decoded = pipe.vae.decode(latent, return_dict=False)[0]
|
||||
image = pipe.image_processor.postprocess(
|
||||
decoded.detach(), output_type="pil")[0]
|
||||
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
|
||||
image.save(OUTPUT_DIR / filename)
|
||||
return {"status": "ok", "filename": filename,
|
||||
"seconds": round(time.monotonic() - started, 3),
|
||||
"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
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
ACTIVE = False
|
||||
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, fmt: str, *args: object) -> None:
|
||||
print(f"[flux9b-beta] {self.client_address[0]} {fmt % args}", flush=True)
|
||||
|
||||
def reply(self, status: int, payload: dict) -> None:
|
||||
body = json.dumps(payload, separators=(",", ":")).encode()
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def do_GET(self) -> None: # noqa: N802
|
||||
if self.path == "/health":
|
||||
self.reply(200, {"status": "ok", "model_loaded": ACTIVE,
|
||||
"model": "FLUX.2-klein-9B-fp8-beta"})
|
||||
else:
|
||||
self.reply(404, {"error": "not found"})
|
||||
|
||||
def do_POST(self) -> None: # noqa: N802
|
||||
if self.headers.get("Authorization", "") != f"Bearer {TOKEN}":
|
||||
self.reply(401, {"error": "unauthorized"})
|
||||
return
|
||||
if self.path != "/generate":
|
||||
self.reply(404, {"error": "not found"})
|
||||
return
|
||||
try:
|
||||
length = int(self.headers.get("Content-Length", "0"))
|
||||
if length < 2 or length > 16384:
|
||||
raise ValueError("invalid request size")
|
||||
self.reply(200, generate(json.loads(self.rfile.read(length))))
|
||||
except Exception as exc:
|
||||
print(f"[flux9b-beta] generation failed: {type(exc).__name__}: "
|
||||
f"{str(exc)[:1000]}", flush=True)
|
||||
self.reply(500, {"status": "error", "message": str(exc)})
|
||||
|
||||
|
||||
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|
||||
Reference in New Issue
Block a user