mcp-tools/FLUX.2-dev
1
1import os2import subprocess3import sys4import io5import gradio as gr6import numpy as np7import random8import spaces9import torch10from diffusers import Flux2Pipeline, Flux2Transformer2DModel11from diffusers import BitsAndBytesConfig as DiffBitsAndBytesConfig12from optimization import optimize_pipeline_13import requests14from PIL import Image15import json16import base6417from huggingface_hub import InferenceClient18 19subprocess.check_call([sys.executable, "-m", "pip", "install", "spaces==0.43.0"])20 21dtype = torch.bfloat1622device = "cuda" if torch.cuda.is_available() else "cpu"23 24MAX_SEED = np.iinfo(np.int32).max25MAX_IMAGE_SIZE = 102426 27# hf_client = InferenceClient(28# api_key=os.environ.get("HF_TOKEN"),29# )30# VLM_MODEL = "baidu/ERNIE-4.5-VL-424B-A47B-Base-PT"31 32# SYSTEM_PROMPT_TEXT_ONLY = """You are an expert prompt engineer for FLUX.2 by Black Forest Labs. Rewrite user prompts to be more descriptive while strictly preserving their core subject and intent.33 34# Guidelines:35# 1. Structure: Keep structured inputs structured (enhance within fields). Convert natural language to detailed paragraphs.36# 2. Details: Add concrete visual specifics - form, scale, textures, materials, lighting (quality, direction, color), shadows, spatial relationships, and environmental context.37# 3. Text in Images: Put ALL text in quotation marks, matching the prompt's language. Always provide explicit quoted text for objects that would contain text in reality (signs, labels, screens, etc.) - without it, the model generates gibberish.38 39# Output only the revised prompt and nothing else."""40 41# SYSTEM_PROMPT_WITH_IMAGES = """You are FLUX.2 by Black Forest Labs, an image-editing expert. You convert editing requests into one concise instruction (50-80 words, ~30 for brief requests).42 43# Rules:44# - Single instruction only, no commentary45# - Use clear, analytical language (avoid "whimsical," "cascading," etc.)46# - Specify what changes AND what stays the same (face, lighting, composition)47# - Reference actual image elements48# - Turn negatives into positives ("don't change X" → "keep X")49# - Make abstractions concrete ("futuristic" → "glowing cyan neon, metallic panels")50# - Keep content PG-1351 52# Output only the final instruction in plain text and nothing else."""53 54def remote_text_encoder(prompts):55 from gradio_client import Client56 57 client = Client("multimodalart/mistral-text-encoder")58 result = client.predict(59 prompt=prompts,60 api_name="/encode_text"61 )62 63 # Load returns a tensor, usually on CPU by default64 prompt_embeds = torch.load(result[0])65 return prompt_embeds66 67# Load model68repo_id = "black-forest-labs/FLUX.2-dev"69 70dit = Flux2Transformer2DModel.from_pretrained(71 repo_id,72 subfolder="transformer",73 torch_dtype=torch.bfloat1674)75 76pipe = Flux2Pipeline.from_pretrained(77 repo_id,78 text_encoder=None,79 transformer=dit,80 torch_dtype=torch.bfloat1681)82pipe.to(device)83 84pipe.transformer.set_attention_backend("_flash_3_hub")85 86# Optimization runs once at startup87optimize_pipeline_(88 pipe,89 image=[Image.new("RGB", (1024, 1024))],90 prompt_embeds = remote_text_encoder("prompt").to(device),91 guidance_scale=2.5,92 width=1024,93 height=1024,94 num_inference_steps=195)96 97def image_to_data_uri(img):98 buffered = io.BytesIO()99 img.save(buffered, format="PNG")100 img_str = base64.b64encode(buffered.getvalue()).decode("utf-8")101 return f"data:image/png;base64,{img_str}"102 103# def upsample_prompt_logic(prompt, image_list):104# try:105# if image_list and len(image_list) > 0:106# # Image + Text Editing Mode107# system_content = SYSTEM_PROMPT_WITH_IMAGES108 109# # Construct user message with text and images110# user_content = [{"type": "text", "text": prompt}]111 112# for img in image_list:113# data_uri = image_to_data_uri(img)114# user_content.append({115# "type": "image_url",116# "image_url": {"url": data_uri}117# })118 119# messages = [120# {"role": "system", "content": system_content},121# {"role": "user", "content": user_content}122# ]123# else:124# # Text Only Mode125# system_content = SYSTEM_PROMPT_TEXT_ONLY126# messages = [127# {"role": "system", "content": system_content},128# {"role": "user", "content": prompt}129# ]130 131# completion = hf_client.chat.completions.create(132# model=VLM_MODEL,133# messages=messages,134# max_tokens=1024135# )136 137# return completion.choices[0].message.content138# except Exception as e:139# print(f"Upsampling failed: {e}")140# return prompt141 142# Updated duration function to match generate_image arguments (including progress)143def get_duration(prompt_embeds, image_list, width, height, num_inference_steps, guidance_scale, seed, force_dimensions, progress=gr.Progress(track_tqdm=True)):144 num_images = 0 if image_list is None else len(image_list)145 step_duration = 1 + 0.8 * num_images146 return max(65, num_inference_steps * step_duration + 10)147 148@spaces.GPU(duration=get_duration)149def generate_image(prompt_embeds, image_list, width, height, num_inference_steps, guidance_scale, seed, force_dimensions, progress=gr.Progress(track_tqdm=True)):150 # Move embeddings to GPU only when inside the GPU decorated function151 prompt_embeds = prompt_embeds.to(device)152 153 generator = torch.Generator(device=device).manual_seed(seed)154 155 pipe_kwargs = {156 "prompt_embeds": prompt_embeds,157 "image": image_list,158 "num_inference_steps": num_inference_steps,159 "guidance_scale": guidance_scale,160 "generator": generator,161 }162 163 if image_list is None or force_dimensions:164 pipe_kwargs["width"] = width165 pipe_kwargs["height"] = height166 167 # Progress bar for the actual generation steps168 if progress:169 progress(0, desc="Starting generation...")170 171 image = pipe(**pipe_kwargs).images[0]172 return image173 174def infer(prompt, input_images=None, seed=42, randomize_seed=False, width=1024, height=1024, num_inference_steps=50, guidance_scale=2.5, force_dimensions=False, prompt_upsampling=False, progress=gr.Progress(track_tqdm=True)):175 176 if randomize_seed:177 seed = random.randint(0, MAX_SEED)178 179 # Prepare image list (convert None or empty gallery to None)180 image_list = None181 if input_images is not None and len(input_images) > 0:182 image_list = []183 for item in input_images:184 image_list.append(item[0])185 186 # 1. Upsampling (Network bound - No GPU needed)187 final_prompt = prompt188 # if prompt_upsampling:189 # progress(0.05, desc="Upsampling prompt...")190 # final_prompt = upsample_prompt_logic(prompt, image_list)191 # print(f"Original Prompt: {prompt}")192 # print(f"Upsampled Prompt: {final_prompt}")193 194 # 2. Text Encoding (Network bound - No GPU needed)195 progress(0.1, desc="Encoding prompt...")196 # This returns CPU tensors197 prompt_embeds = remote_text_encoder(final_prompt)198 199 # 3. Image Generation (GPU bound)200 progress(0.3, desc="Waiting for GPU...")201 image = generate_image(202 prompt_embeds, 203 image_list, 204 width, 205 height, 206 num_inference_steps, 207 guidance_scale, 208 seed, 209 force_dimensions,210 progress211 )212 213 return image, "Seed used for generation: " + str(seed)214 215examples = [216 ["Create a vase on a table in living room, the color of the vase is a gradient of color, starting with #02eb3c color and finishing with #edfa3c. The flowers inside the vase have the color #ff0088"],217 ["Photorealistic infographic showing the complete Berlin TV Tower (Fernsehturm) from ground base to antenna tip, full vertical view with entire structure visible including concrete shaft, metallic sphere, and antenna spire. Slight upward perspective angle looking up toward the iconic sphere, perfectly centered on clean white background. Left side labels with thin horizontal connector lines: the text '368m' in extra large bold dark grey numerals (#2D3748) positioned at exactly the antenna tip with 'TOTAL HEIGHT' in small caps below. The text '207m' in extra large bold with 'TELECAFÉ' in small caps below, with connector line touching the sphere precisely at the window level. Right side label with horizontal connector line touching the sphere's equator: the text '32m' in extra large bold dark grey numerals with 'SPHERE DIAMETER' in small caps below. Bottom section arranged in three balanced columns: Left - Large text '986' in extra bold dark grey with 'STEPS' in caps below. Center - 'BERLIN TV TOWER' in bold caps with 'FERNSEHTURM' in lighter weight below. Right - 'INAUGURATED' in bold caps with 'OCTOBER 3, 1969' below. All typography in modern sans-serif font (such as Inter or Helvetica), color #2D3748, clean minimal technical diagram style. Horizontal connector lines are thin, precise, and clearly visible, touching the tower structure at exact corresponding measurement points. Professional architectural elevation drawing aesthetic with dynamic low angle perspective creating sense of height and grandeur, poster-ready infographic design with perfect visual hierarchy."],218 ["Soaking wet capybara taking shelter under a banana leaf in the rainy jungle, close up photo"],219 ["A kawaii die-cut sticker of a chubby orange cat, featuring big sparkly eyes and a happy smile with paws raised in greeting and a heart-shaped pink nose. The design should have smooth rounded lines with black outlines and soft gradient shading with pink cheeks."],220]221 222examples_images = [223 # ["Replace the top of the person from image 1 with the one from image 2", ["person1.webp", "woman2.webp"]],224 ["The person from image 1 is petting the cat from image 2, the bird from image 3 is next to them", ["woman1.webp", "cat_window.webp", "bird.webp"]]225]226 227css="""228#col-container {229 margin: 0 auto;230 max-width: 620px;231}232.gallery-container img{233 object-fit: contain;234}235"""236 237with gr.Blocks() as demo:238 239 with gr.Column(elem_id="col-container"):240 gr.Markdown(f"""# FLUX.2 [dev]241FLUX.2 [dev] is a 32B model rectified flow capable of generating, editing and combining images based on text instructions model [[model](https://huggingface.co/black-forest-labs/FLUX.2-dev)], [[blog](https://bfl.ai/blog/flux-2)]242 """)243 244 with gr.Accordion("Input image(s) (optional)", open=True):245 input_images = gr.Gallery(246 label="Input Image(s)",247 type="pil",248 columns=3,249 rows=1,250 )251 252 with gr.Row():253 254 prompt = gr.Text(255 label="Prompt",256 show_label=False,257 max_lines=2,258 placeholder="Enter your prompt",259 container=False,260 scale=3261 )262 263 run_button = gr.Button("Run", scale=1)264 265 result = gr.Image(label="Result", show_label=False)266 267 with gr.Accordion("Advanced Settings", open=False):268 269 prompt_upsampling = gr.Checkbox(270 label="Prompt Upsampling",271 value=True,272 info="Automatically enhance the prompt using a VLM"273 )274 275 seed = gr.Slider(276 label="Seed",277 minimum=0,278 maximum=MAX_SEED,279 step=1,280 value=0,281 )282 283 randomize_seed = gr.Checkbox(label="Randomize seed", value=True)284 285 with gr.Row():286 287 width = gr.Slider(288 label="Width",289 minimum=256,290 maximum=MAX_IMAGE_SIZE,291 step=32,292 value=1024,293 )294 295 height = gr.Slider(296 label="Height",297 minimum=256,298 maximum=MAX_IMAGE_SIZE,299 step=32,300 value=1024,301 )302 303 force_dimensions = gr.Checkbox(304 label="Force width/height when image input",305 value=False,306 info="When unchecked, width/height settings are ignored if input images are provided"307 )308 309 with gr.Row():310 311 num_inference_steps = gr.Slider(312 label="Number of inference steps",313 minimum=1,314 maximum=100,315 step=1,316 value=30,317 )318 319 guidance_scale = gr.Slider(320 label="Guidance scale",321 minimum=0.0,322 maximum=10.0,323 step=0.1,324 value=4,325 )326 327 gr.Examples(328 examples=examples,329 fn=infer,330 inputs=[prompt],331 outputs=[result, seed],332 cache_examples=True,333 cache_mode="lazy"334 )335 336 gr.Examples(337 examples=examples_images,338 fn=infer,339 inputs=[prompt, input_images],340 outputs=[result, seed],341 cache_examples=True,342 cache_mode="lazy"343 )344 345 gr.on(346 triggers=[run_button.click, prompt.submit],347 fn=infer,348 inputs=[prompt, input_images, seed, randomize_seed, width, height, num_inference_steps, guidance_scale, force_dimensions, prompt_upsampling],349 outputs=[result, seed]350 )351 352demo.launch(mcp_server=True,css=css)