diffusers/flux-quant
9
1import torch2import gradio as gr3from diffusers import FluxPipeline, FluxTransformer2DModel4import gc5import random6import glob7from pathlib import Path8from PIL import Image9import os10import time11import json12from fasteners import InterProcessLock13import spaces14from datasets import Dataset, Image as HFImage, load_dataset15from datasets import Features, Value16from datasets import concatenate_datasets17from datetime import datetime18 19AGG_FILE = Path(__file__).parent / "agg_stats.json"20LOCK_FILE = AGG_FILE.with_suffix(".lock")21 22def _load_agg_stats() -> dict:23 if AGG_FILE.exists():24 with open(AGG_FILE, "r") as f:25 try:26 return json.load(f)27 except json.JSONDecodeError:28 print(f"Warning: {AGG_FILE} is corrupted. Starting with empty stats.")29 return {"8-bit bnb": {"attempts": 0, "correct": 0}, "4-bit bnb": {"attempts": 0, "correct": 0}}30 return {"8-bit bnb": {"attempts": 157, "correct": 74},31 "4-bit bnb": {"attempts": 159, "correct": 78}}32 33def _save_agg_stats(stats: dict) -> None:34 with InterProcessLock(str(LOCK_FILE)):35 with open(AGG_FILE, "w") as f:36 json.dump(stats, f, indent=2)37 38DEVICE = "cuda" if torch.cuda.is_available() else "cpu"39print(f"Using device: {DEVICE}")40 41DEFAULT_HEIGHT = 102442DEFAULT_WIDTH = 102443DEFAULT_GUIDANCE_SCALE = 3.544DEFAULT_NUM_INFERENCE_STEPS = 1545DEFAULT_MAX_SEQUENCE_LENGTH = 51246HF_TOKEN = os.environ.get("HF_ACCESS_TOKEN")47HF_DATASET_REPO_ID = "diffusers/flux-quant-challenge-submissions"48 49CACHED_PIPES = {}50def load_bf16_pipeline():51 print("Loading BF16 pipeline...")52 MODEL_ID = "black-forest-labs/FLUX.1-dev"53 if MODEL_ID in CACHED_PIPES:54 return CACHED_PIPES[MODEL_ID]55 start_time = time.time()56 try:57 pipe = FluxPipeline.from_pretrained(58 MODEL_ID,59 torch_dtype=torch.bfloat16,60 token=HF_TOKEN61 )62 # pipe.to(DEVICE)63 pipe.enable_model_cpu_offload()64 end_time = time.time()65 mem_reserved = torch.cuda.memory_reserved(0)/1024**3 if DEVICE == "cuda" else 066 print(f"BF16 Pipeline loaded in {end_time - start_time:.2f}s. Memory reserved: {mem_reserved:.2f} GB")67 CACHED_PIPES[MODEL_ID] = pipe68 return pipe69 except Exception as e:70 print(f"Error loading BF16 pipeline: {e}")71 raise72 73def load_bnb_8bit_pipeline():74 print("Loading 8-bit BNB pipeline...")75 MODEL_ID = "derekl35/FLUX.1-dev-bnb-8bit"76 if MODEL_ID in CACHED_PIPES:77 return CACHED_PIPES[MODEL_ID]78 start_time = time.time()79 try:80 pipe = FluxPipeline.from_pretrained(81 MODEL_ID,82 torch_dtype=torch.bfloat1683 )84 # pipe.to(DEVICE)85 pipe.enable_model_cpu_offload()86 end_time = time.time()87 mem_reserved = torch.cuda.memory_reserved(0)/1024**3 if DEVICE == "cuda" else 088 print(f"8-bit BNB pipeline loaded in {end_time - start_time:.2f}s. Memory reserved: {mem_reserved:.2f} GB")89 CACHED_PIPES[MODEL_ID] = pipe90 return pipe91 except Exception as e:92 print(f"Error loading 8-bit BNB pipeline: {e}")93 raise94 95def load_bnb_4bit_pipeline():96 print("Loading 4-bit BNB pipeline...")97 MODEL_ID = "derekl35/FLUX.1-dev-nf4"98 if MODEL_ID in CACHED_PIPES:99 return CACHED_PIPES[MODEL_ID]100 start_time = time.time()101 try:102 pipe = FluxPipeline.from_pretrained(103 MODEL_ID,104 torch_dtype=torch.bfloat16105 )106 # pipe.to(DEVICE)107 pipe.enable_model_cpu_offload()108 end_time = time.time()109 mem_reserved = torch.cuda.memory_reserved(0)/1024**3 if DEVICE == "cuda" else 0110 print(f"4-bit BNB pipeline loaded in {end_time - start_time:.2f}s. Memory reserved: {mem_reserved:.2f} GB")111 CACHED_PIPES[MODEL_ID] = pipe112 return pipe113 except Exception as e:114 print(f"Error loading 4-bit BNB pipeline: {e}")115 raise116 117@spaces.GPU(duration=240)118def generate_images(prompt, quantization_choice, progress=gr.Progress(track_tqdm=True)):119 if not prompt:120 return None, {}, gr.update(value="Please enter a prompt.", interactive=False), None, [], gr.update(interactive=True), gr.update(interactive=True)121 122 if not quantization_choice:123 return None, {}, gr.update(value="Please select a quantization method.", interactive=False), None, [], gr.update(interactive=True), gr.update(interactive=True)124 125 if quantization_choice == "8-bit bnb":126 quantized_load_func = load_bnb_8bit_pipeline127 quantized_label = "Quantized (8-bit bnb)"128 elif quantization_choice == "4-bit bnb":129 quantized_load_func = load_bnb_4bit_pipeline130 quantized_label = "Quantized (4-bit bnb)"131 else:132 return None, {}, gr.update(value="Invalid quantization choice.", interactive=False), None, [], gr.update(interactive=True), gr.update(interactive=True)133 134 model_configs = [135 ("Original", load_bf16_pipeline),136 (quantized_label, quantized_load_func),137 ]138 139 results = []140 pipe_kwargs = {141 "prompt": prompt,142 "height": DEFAULT_HEIGHT,143 "width": DEFAULT_WIDTH,144 "guidance_scale": DEFAULT_GUIDANCE_SCALE,145 "num_inference_steps": DEFAULT_NUM_INFERENCE_STEPS,146 "max_sequence_length": DEFAULT_MAX_SEQUENCE_LENGTH,147 }148 149 seed = random.getrandbits(64)150 print(f"Using seed: {seed}")151 152 for i, (label, load_func) in enumerate(model_configs):153 progress(i / len(model_configs), desc=f"Loading {label} model...")154 print(f"\n--- Loading {label} Model ---")155 load_start_time = time.time()156 try:157 current_pipe = load_func()158 load_end_time = time.time()159 print(f"{label} model loaded in {load_end_time - load_start_time:.2f} seconds.")160 161 progress((i + 0.5) / len(model_configs), desc=f"Generating with {label} model...")162 print(f"--- Generating with {label} Model ---")163 gen_start_time = time.time()164 image_list = current_pipe(**pipe_kwargs, generator=torch.manual_seed(seed)).images165 image = image_list[0]166 gen_end_time = time.time()167 results.append({"label": label, "image": image})168 print(f"--- Finished Generation with {label} Model in {gen_end_time - gen_start_time:.2f} seconds ---")169 mem_reserved = torch.cuda.memory_reserved(0)/1024**3 if DEVICE == "cuda" else 0170 print(f"Memory reserved: {mem_reserved:.2f} GB")171 172 except Exception as e:173 print(f"Error during {label} model processing: {e}")174 return None, {}, gr.update(value=f"Error processing {label} model: {e}", interactive=False), None, [], gr.update(interactive=True), gr.update(interactive=True)175 176 177 if len(results) != len(model_configs):178 return None, {}, gr.update(value="Failed to generate images for all model types.", interactive=False), None, [], gr.update(interactive=True), gr.update(interactive=True)179 180 shuffled_results = results.copy()181 random.shuffle(shuffled_results)182 shuffled_data_for_gallery = [(res["image"], f"Image {i+1}") for i, res in enumerate(shuffled_results)]183 correct_mapping = {i: res["label"] for i, res in enumerate(shuffled_results)}184 print("Correct mapping (hidden):", correct_mapping)185 186 return shuffled_data_for_gallery, correct_mapping, prompt, seed, results, "Generation complete! Make your guess.", None, gr.update(interactive=True), gr.update(interactive=True)187 188 189def check_guess(user_guess, correct_mapping_state):190 if not isinstance(correct_mapping_state, dict) or not correct_mapping_state:191 return "Please generate images first (state is empty or invalid)."192 if user_guess is None:193 return "Please select which image you think is quantized."194 195 quantized_image_index = -1196 quantized_label_actual = ""197 for index, label in correct_mapping_state.items():198 if "Quantized" in label:199 quantized_image_index = index200 quantized_label_actual = label201 break202 if quantized_image_index == -1:203 return "Error: Could not find the quantized image in the mapping data."204 205 correct_guess_label = f"Image {quantized_image_index + 1}"206 if user_guess == correct_guess_label:207 feedback = f"Correct! {correct_guess_label} used the {quantized_label_actual} model."208 else:209 feedback = f"Incorrect. The quantized image ({quantized_label_actual}) was {correct_guess_label}."210 return feedback211 212EXAMPLE_DIR = Path(__file__).parent / "examples"213EXAMPLES = [214 {215 "prompt": "A photorealistic portrait of an astronaut on Mars",216 "files": ["astronauts_seed_6456306350371904162.png", "astronauts_bnb_8bit.png"],217 "quantized_idx": 1,218 "quant_method": "8-bit bnb",219 "summary": "Astronaut on Mars",220 },221 {222 "prompt": "Water-color painting of a cat wearing sunglasses",223 "files": ["watercolor_cat_bnb_8bit.png", "watercolor_cat_seed_14269059182221286790.png"],224 "quantized_idx": 0,225 "quant_method": "8-bit bnb",226 "summary": "Cat with Sunglasses",227 },228 # {229 # "prompt": "Neo-tokyo cyberpunk cityscape at night, rain-soaked streets, 8-K",230 # "files": ["cyber_city_q.jpg", "cyber_city.jpg"],231 # "quantized_idx": 0,232 # },233]234 235def load_example(idx):236 ex = EXAMPLES[idx]237 imgs = [Image.open(EXAMPLE_DIR / f) for f in ex["files"]]238 gallery_items = [(img, f"Image {i+1}") for i, img in enumerate(imgs)]239 mapping = {i: (f"Quantized ({ex['quant_method']})" if i == ex["quantized_idx"] else "Original")240 for i in range(2)}241 return gallery_items, mapping, f"{ex['prompt']}"242 243def _accuracy_string(correct: int, attempts: int) -> tuple[str, float]:244 if attempts:245 pct = 100 * correct / attempts246 return f"{pct:.1f}%", pct247 return "N/A", -1.0248 249def update_leaderboards_data():250 agg = _load_agg_stats()251 quant_rows = []252 for method, stats in agg.items():253 acc_str, acc_val = _accuracy_string(stats["correct"], stats["attempts"])254 quant_rows.append([255 method,256 stats["correct"],257 stats["attempts"],258 acc_str259 ])260 quant_rows.sort(key=lambda r: r[1]/r[2] if r[2] != 0 else 1e9)261 return quant_rows262 263quant_df = gr.DataFrame(264 headers=["Method", "Correct Guesses", "Total Attempts", "Detectability %"],265 interactive=False, col_count=(4, "fixed")266)267 268with gr.Blocks(title="FLUX Quantization Challenge", theme=gr.themes.Soft()) as demo:269 gr.Markdown("# FLUX Model Quantization Challenge")270 with gr.Tabs():271 with gr.TabItem("Challenge"):272 gr.Markdown(273 "Compare the original FLUX.1-dev (BF16) model against a quantized version (4-bit or 8-bit bnb). "274 "Enter a prompt, choose the quantization method, and generate two images. "275 "The images will be shuffled, can you spot which one was quantized?"276 )277 278 gr.Markdown("### Examples")279 ex_selector = gr.Radio(280 choices=[ex["summary"] for ex in EXAMPLES],281 label="Choose an example prompt",282 interactive=True,283 )284 gr.Markdown("### …or create your own comparison")285 with gr.Row():286 prompt_input = gr.Textbox(label="Enter Prompt", scale=3)287 quantization_choice_radio = gr.Radio(288 choices=["8-bit bnb", "4-bit bnb"],289 label="Select Quantization",290 value="8-bit bnb",291 scale=1292 )293 generate_button = gr.Button("Generate & Compare", variant="primary", scale=1)294 295 output_gallery = gr.Gallery(296 label="Generated Images",297 columns=2,298 height=606,299 object_fit="contain",300 allow_preview=True,301 show_label=True,302 )303 304 gr.Markdown("### Which image used the selected quantization method?")305 with gr.Row():306 image1_btn = gr.Button("Image 1")307 image2_btn = gr.Button("Image 2")308 309 feedback_box = gr.Textbox(label="Feedback", interactive=False, lines=1)310 311 with gr.Row():312 session_score_box = gr.Textbox(label="Your accuracy this session", interactive=False)313 314 gr.Markdown("""315 ### Dataset Information316 Unless you opt out below, your submissions will be recorded in a dataset. This dataset contains anonymized challenge results including prompts, images, quantization methods, 317 and whether guesses were correct.318 """)319 320 opt_out_checkbox = gr.Checkbox(321 label="Opt out of data collection (don't record my submissions to the dataset)", 322 value=False323 )324 325 correct_mapping_state = gr.State({})326 session_stats_state = gr.State(327 {"8-bit bnb": {"attempts": 0, "correct": 0},328 "4-bit bnb": {"attempts": 0, "correct": 0}}329 )330 is_example_state = gr.State(False)331 prompt_state = gr.State("")332 seed_state = gr.State(None)333 results_state = gr.State([])334 335 def _load_example_and_update_dfs(sel_summary):336 idx = next((i for i, ex in enumerate(EXAMPLES) if ex["summary"] == sel_summary), -1)337 if idx == -1:338 print(f"Error: Example with summary '{sel_summary}' not found.")339 return (gr.update(), gr.update(), gr.update(), False, gr.update(), "", None, [])340 341 ex = EXAMPLES[idx]342 gallery_items, mapping, prompt = load_example(idx)343 quant_data = update_leaderboards_data()344 return gallery_items, mapping, prompt, True, quant_data, "", None, []345 346 ex_selector.change(347 fn=_load_example_and_update_dfs,348 inputs=ex_selector,349 outputs=[output_gallery, correct_mapping_state, prompt_input, is_example_state, quant_df,350 prompt_state, seed_state, results_state],351 ).then(352 lambda: (gr.update(interactive=True), gr.update(interactive=True)),353 outputs=[image1_btn, image2_btn],354 )355 356 generate_button.click(357 fn=generate_images,358 inputs=[prompt_input, quantization_choice_radio],359 outputs=[output_gallery, correct_mapping_state, prompt_state, seed_state, results_state,360 feedback_box]361 ).then(362 lambda: False, # for is_example_state363 outputs=[is_example_state]364 ).then(365 lambda: (gr.update(interactive=True),366 gr.update(interactive=True),367 ""),368 outputs=[image1_btn, image2_btn, feedback_box],369 )370 371 def choose(choice_string, mapping, session_stats, is_example,372 prompt, seed, results, opt_out):373 feedback = check_guess(choice_string, mapping)374 375 if not mapping:376 return feedback, gr.update(), gr.update(), "", session_stats, gr.update()377 378 quant_label_from_mapping = next((label for label in mapping.values() if "Quantized" in label), None)379 if not quant_label_from_mapping:380 print("Error: Could not determine quantization label from mapping:", mapping)381 return ("Internal Error: Could not process results.", gr.update(interactive=False), gr.update(interactive=False),382 "", session_stats, gr.update())383 384 quant_key = "8-bit bnb" if "8-bit bnb" in quant_label_from_mapping else "4-bit bnb"385 got_it_right = "Correct!" in feedback386 sess = session_stats.copy()387 388 if not is_example: # Only log and update stats if it's not an example run389 sess[quant_key]["attempts"] += 1390 if got_it_right:391 sess[quant_key]["correct"] += 1392 session_stats = sess # Update the state for the UI393 394 AGG_STATS = _load_agg_stats()395 AGG_STATS[quant_key]["attempts"] += 1396 if got_it_right:397 AGG_STATS[quant_key]["correct"] += 1398 _save_agg_stats(AGG_STATS)399 400 if not HF_TOKEN:401 print("Warning: HF_TOKEN not set. Skipping dataset logging.")402 elif not results:403 print("Warning: Results state is empty. Skipping dataset logging.")404 elif opt_out:405 print("User opted out of dataset logging. Skipping.")406 else:407 print(f"Logging guess to HF Dataset: {HF_DATASET_REPO_ID}")408 original_image = None409 quantized_image = None410 quantized_image_pos = -1411 412 for shuffled_idx, original_label in mapping.items():413 if "Quantized" in original_label:414 quantized_image_pos = shuffled_idx415 break416 417 original_image = next((res["image"] for res in results if "Original" in res["label"]), None)418 quantized_image = next((res["image"] for res in results if "Quantized" in res["label"]), None)419 420 if original_image and quantized_image:421 expected_features = Features({422 "timestamp": Value("string"),423 "prompt": Value("string"),424 "quantization_method": Value("string"),425 "seed": Value("string"),426 "image_original": HFImage(),427 "image_quantized": HFImage(),428 "quantized_image_displayed_position": Value("string"),429 "user_guess_displayed_position": Value("string"),430 "correct_guess": Value("bool"),431 "username": Value("string"), # Handles None432 })433 434 new_data_dict_of_lists = {435 "timestamp": [datetime.now().isoformat()],436 "prompt": [prompt],437 "quantization_method": [quant_key],438 "seed": [str(seed)],439 "image_original": [original_image],440 "image_quantized": [quantized_image],441 "quantized_image_displayed_position": [f"Image {quantized_image_pos + 1}"],442 "user_guess_displayed_position": [choice_string],443 "correct_guess": [got_it_right],444 "username": [None], # Log None for username445 }446 try:447 existing_ds = load_dataset(448 HF_DATASET_REPO_ID,449 split="train",450 token=HF_TOKEN,451 features=expected_features,452 )453 new_row_ds = Dataset.from_dict(new_data_dict_of_lists, features=expected_features)454 combined_ds = concatenate_datasets([existing_ds, new_row_ds])455 combined_ds.push_to_hub(HF_DATASET_REPO_ID, token=HF_TOKEN, split="train")456 print(f"Successfully appended guess to {HF_DATASET_REPO_ID} (train split)")457 except Exception as e:458 print(f"Could not load or append to existing dataset/split. Creating 'train' split with the new item. Error: {e}")459 ds_new = Dataset.from_dict(new_data_dict_of_lists, features=expected_features)460 ds_new.push_to_hub(HF_DATASET_REPO_ID, token=HF_TOKEN, split="train")461 print(f"Successfully created and logged new 'train' split to {HF_DATASET_REPO_ID}")462 else:463 print("Error: Could not find original or quantized image in results state for logging.")464 465 def _fmt(d):466 a, c = d["attempts"], d["correct"]467 pct = 100 * c / a if a else 0468 return f"{c} / {a} ({pct:.1f}%)"469 470 session_msg = ", ".join(471 f"{k}: {_fmt(v)}" for k, v in sess.items()472 )473 474 quant_data = update_leaderboards_data()475 return (feedback,476 gr.update(interactive=False),477 gr.update(interactive=False),478 session_msg,479 session_stats, # Return the potentially updated session_stats480 quant_data)481 482 image1_btn.click(483 fn=lambda mapping, sess, is_ex, p, s, r, opt_out: choose("Image 1", mapping, sess, is_ex, p, s, r, opt_out),484 inputs=[correct_mapping_state, session_stats_state, is_example_state,485 prompt_state, seed_state, results_state, opt_out_checkbox],486 outputs=[feedback_box, image1_btn, image2_btn,487 session_score_box, session_stats_state,488 quant_df],489 )490 image2_btn.click(491 fn=lambda mapping, sess, is_ex, p, s, r, opt_out: choose("Image 2", mapping, sess, is_ex, p, s, r, opt_out),492 inputs=[correct_mapping_state, session_stats_state, is_example_state,493 prompt_state, seed_state, results_state, opt_out_checkbox],494 outputs=[feedback_box, image1_btn, image2_btn,495 session_score_box, session_stats_state,496 quant_df],497 )498 499 with gr.TabItem("Leaderboard"):500 gr.Markdown("## Quantization Method Leaderboard *(Lower % ⇒ harder to detect)*")501 leaderboard_tab_quant_df = gr.DataFrame(502 headers=["Method", "Correct Guesses", "Total Attempts", "Detectability %"],503 interactive=False, col_count=(4, "fixed"), label="Quantization Method Leaderboard"504 )505 506 def update_all_leaderboards_for_tab():507 q_rows = update_leaderboards_data()508 return q_rows # Only return quantization method data509 510 demo.load(update_all_leaderboards_for_tab, outputs=[511 leaderboard_tab_quant_df,512 ])513 514if __name__ == "__main__":515 demo.launch(share=True)