Files

313 lines
12 KiB
Python

#!/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 contextlib import contextmanager
from threading import Lock
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
GENERATION_LOCK = Lock()
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}"))
@contextmanager
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
try:
yield scales
finally:
single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = original
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
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
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()
transformer = text_encoder = pipe = latent = decoded = image = None
prompt_embeds = generator = module = None
kwargs = {}
scales = {}
try:
enable_huggingface_checkpointing()
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,
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")
module = None
scales.clear()
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
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)
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:
# 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()
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)})
if __name__ == "__main__":
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()