#!/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()