import os import random import traceback from functools import lru_cache import gradio as gr import torch from PIL import Image from diffusers import Flux2KleinPipeline MODEL_ID = "black-forest-labs/FLUX.2-klein-4b-fp8" # FLUX.2 Klein is intended for GPU use. The FP8 checkpoint is about 4.07 GB, # while the full 4B model is much larger. DEVICE = "cuda" if torch.cuda.is_available() else "cpu" # Prefer bfloat16 for the surrounding pipeline components. The transformer # checkpoint itself is FP8 and is loaded from the model repository. DTYPE = torch.bfloat16 _pipe = None def load_pipeline(): global _pipe if _pipe is not None: return _pipe if DEVICE != "cuda": raise RuntimeError( "This Space requires a CUDA GPU for practical FLUX.2 Klein inference. " "Please enable a GPU hardware accelerator in the Space settings." ) # from_pretrained() handles the repository's model configuration and FP8 # transformer weights. _pipe = Flux2KleinPipeline.from_pretrained( MODEL_ID, torch_dtype=DTYPE, ) # CPU offload reduces peak VRAM usage and is useful on smaller GPU flavors. _pipe.enable_model_cpu_offload() return _pipe def make_generator(seed: int): # The generator must be created on the CUDA device for deterministic CUDA # sampling. return torch.Generator(device="cuda").manual_seed(int(seed)) def generate( prompt, reference_images, width, height, steps, guidance, seed, ): if not prompt or not prompt.strip(): raise gr.Error("Enter a prompt first.") if DEVICE != "cuda": raise gr.Error( "No CUDA GPU is available. Enable a GPU accelerator for this Space." ) try: pipe = load_pipeline() seed = int(seed) if seed < 0: seed = random.randint(0, 2**31 - 1) # Gradio can return a single PIL image or a list depending on the # component configuration. images = [] if reference_images: if isinstance(reference_images, Image.Image): images = [reference_images] else: images = [ item for item in reference_images if isinstance(item, Image.Image) ] kwargs = dict( prompt=prompt.strip(), height=int(height), width=int(width), guidance_scale=float(guidance), num_inference_steps=int(steps), generator=make_generator(seed), ) # FLUX.2 Klein supports image-to-image/reference-image editing. # Pass the reference image when one was supplied. if images: kwargs["image"] = images with torch.inference_mode(): result = pipe(**kwargs) image = result.images[0] # Keep the generated image in a normal PIL format for Gradio. if image.mode not in ("RGB", "RGBA"): image = image.convert("RGB") return image, seed except Exception as exc: traceback.print_exc() raise gr.Error(f"Generation failed: {exc}") def clear_reference(): return None with gr.Blocks( title="FLUX.2 Klein 4B FP8", theme=gr.themes.Soft(), ) as demo: gr.Markdown( """ # FLUX.2 [klein] 4B FP8 Fast text-to-image generation and reference-image editing using `black-forest-labs/FLUX.2-klein-4b-fp8`. **Tip:** Leave Reference Images empty for text-to-image generation. Upload one or more images when you want to edit/use them as references. """ ) with gr.Row(): with gr.Column(scale=1): prompt = gr.Textbox( label="Prompt", placeholder=( "A cinematic portrait of a futuristic warrior, " "dramatic rim lighting, detailed concept art" ), lines=5, ) reference_images = gr.Gallery( label="Reference Images (optional)", columns=3, rows=2, height="auto", allow_preview=True, type="pil", ) with gr.Row(): width = gr.Dropdown( choices=[512, 640, 768, 896, 1024, 1152, 1280], value=1024, label="Width", ) height = gr.Dropdown( choices=[512, 640, 768, 896, 1024, 1152, 1280], value=1024, label="Height", ) with gr.Row(): steps = gr.Slider( minimum=1, maximum=12, value=4, step=1, label="Inference Steps", ) guidance = gr.Slider( minimum=0, maximum=10, value=4.0, step=0.5, label="Guidance Scale", ) seed = gr.Number( value=-1, precision=0, label="Seed (-1 = random)", ) with gr.Row(): generate_btn = gr.Button( "Generate", variant="primary", ) clear_btn = gr.Button("Clear Reference") with gr.Column(scale=1): output = gr.Image( label="Generated Image", type="pil", format="png", ) used_seed = gr.Number( label="Used Seed", interactive=False, ) gr.Markdown( """ ### Examples - `A small red panda astronaut exploring a neon-lit alien city, cinematic photography` - `Turn the reference character into a clean 2D animation style, preserve the pose and costume design` - `Transform the reference image into a dramatic fantasy movie poster, detailed lighting, high contrast` """ ) generate_btn.click( fn=generate, inputs=[ prompt, reference_images, width, height, steps, guidance, seed, ], outputs=[output, used_seed], ) clear_btn.click( fn=clear_reference, inputs=[], outputs=[reference_images], ) if __name__ == "__main__": demo.queue(max_size=20).launch()