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