Team Ai
Apppublic

diffusers/flux-quant

sourceHugging Faceupdated 1y agoView on Hugging Face
9likes
app.py515 linesDownload Raw Back to root
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)