Team Ai
Apppublic

michaelcreatesstuff/llm-grounded-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
2likes
app.py310 linesDownload Raw Back to root
1import gradio as gr2import numpy as np3import ast4from matplotlib.patches import Polygon5from matplotlib.collections import PatchCollection6import matplotlib.pyplot as plt7from utils.parse import filter_boxes8from generation import run as run_ours9from baseline import run as run_baseline10import torch11from shared import DEFAULT_SO_NEGATIVE_PROMPT, DEFAULT_OVERALL_NEGATIVE_PROMPT12from examples import stage1_examples, stage2_examples13import base6414import io15 16print(f"Is CUDA available: {torch.cuda.is_available()}")17if torch.cuda.is_available():18    print(f"CUDA device: {torch.cuda.get_device_name(torch.cuda.current_device())}")19 20box_scale = (512, 512)21size = box_scale22 23bg_prompt_text = "Background prompt: "24 25default_template = """You are an intelligent bounding box generator. I will provide you with a caption for a photo, image, or painting. Your task is to generate the bounding boxes for the objects mentioned in the caption, along with a background prompt describing the scene. The images are of size 512x512, and the bounding boxes should not overlap or go beyond the image boundaries. Each bounding box should be in the format of (object name, [top-left x coordinate, top-left y coordinate, box width, box height]) and include exactly one object. Make the boxes larger if possible. Do not put objects that are already provided in the bounding boxes into the background prompt. If needed, you can make reasonable guesses. Generate the object descriptions and background prompts in English even if the caption might not be in English. Do not include non-existing or excluded objects in the background prompt. Please refer to the example below for the desired format.26 27Caption: A realistic image of landscape scene depicting a green car parking on the left of a blue truck, with a red air balloon and a bird in the sky28Objects: [('a green car', [21, 181, 211, 159]), ('a blue truck', [269, 181, 209, 160]), ('a red air balloon', [66, 8, 145, 135]), ('a bird', [296, 42, 143, 100])]29Background prompt: A realistic image of a landscape scene30 31Caption: A watercolor painting of a wooden table in the living room with an apple on it32Objects: [('a wooden table', [65, 243, 344, 206]), ('a apple', [206, 306, 81, 69])]33Background prompt: A watercolor painting of a living room34 35Caption: A watercolor painting of two pandas eating bamboo in a forest36Objects: [('a panda eating bambooo', [30, 171, 212, 226]), ('a panda eating bambooo', [264, 173, 222, 221])]37Background prompt: A watercolor painting of a forest38 39Caption: A realistic image of four skiers standing in a line on the snow near a palm tree40Objects: [('a skier', [5, 152, 139, 168]), ('a skier', [278, 192, 121, 158]), ('a skier', [148, 173, 124, 155]), ('a palm tree', [404, 180, 103, 180])]41Background prompt: A realistic image of an outdoor scene with snow42 43Caption: An oil painting of a pink dolphin jumping on the left of a steam boat on the sea44Objects: [('a steam boat', [232, 225, 257, 149]), ('a jumping pink dolphin', [21, 249, 189, 123])]45Background prompt: An oil painting of the sea46 47Caption: A realistic image of a cat playing with a dog in a park with flowers48Objects: [('a playful cat', [51, 67, 271, 324]), ('a playful dog', [302, 119, 211, 228])]49Background prompt: A realistic image of a park with flowers50 51Caption: 一个客厅场景的油画,墙上挂着电视,电视下面是一个柜子,柜子上有一个花瓶。52Objects: [('a tv', [88, 85, 335, 203]), ('a cabinet', [57, 308, 404, 201]), ('a flower vase', [166, 222, 92, 108])]53Background prompt: An oil painting of a living room scene"""54 55simplified_prompt = """{template}56 57Caption: {prompt}58Objects: """59 60prompt_placeholder = "A realistic photo of a gray cat and an orange dog on the grass."61 62layout_placeholder = """Caption: A realistic photo of a gray cat and an orange dog on the grass.63Objects: [('a gray cat', [67, 243, 120, 126]), ('an orange dog', [265, 193, 190, 210])]64Background prompt: A realistic photo of a grassy area."""65 66canvasbase64 = ""67oursimagebase64 = ""68 69def get_lmd_prompt(prompt, template=default_template):70    if prompt == "":71        prompt = prompt_placeholder72    if template == "":73        template = default_template74    return simplified_prompt.format(template=template, prompt=prompt)75 76def get_layout_image(response):77    global canvasbase6478    if response == "":79        response = layout_placeholder80    gen_boxes, bg_prompt = parse_input(response)81    fig = plt.figure(figsize=(8, 8))82    # https://stackoverflow.com/questions/7821518/save-plot-to-numpy-array83    show_boxes(gen_boxes, bg_prompt)84    # If we haven't already shown or saved the plot, then we need to85    # draw the figure first...86    fig.canvas.draw()87 88    # Now we can save it to a numpy array.89    data = np.frombuffer(fig.canvas.tostring_rgb(), dtype=np.uint8)90    data = data.reshape(fig.canvas.get_width_height()[::-1] + (3,))91    pic_IObytes = io.BytesIO()92    plt.savefig(pic_IObytes,  format='png')93    pic_IObytes.seek(0)94    canvasbase64 = base64.b64encode(pic_IObytes.read()).decode()95  96    plt.clf()97    return [data,canvasbase64]98 99def get_layout_image_gallery(response):100    return get_layout_image(response)101 102def get_ours_image(response, overall_prompt_override="", seed=0, num_inference_steps=20, dpm_scheduler=True, use_autocast=False, fg_seed_start=20, fg_blending_ratio=0.1, frozen_step_ratio=0.4, gligen_scheduled_sampling_beta=0.3, so_negative_prompt=DEFAULT_SO_NEGATIVE_PROMPT, overall_negative_prompt=DEFAULT_OVERALL_NEGATIVE_PROMPT, show_so_imgs=False, scale_boxes=False):103    global oursimagebase64104    if response == "":105        response = layout_placeholder106    gen_boxes, bg_prompt = parse_input(response)107    gen_boxes = filter_boxes(gen_boxes, scale_boxes=scale_boxes)108    spec = {109        # prompt is unused110        'prompt': '',111        'gen_boxes': gen_boxes,112        'bg_prompt': bg_prompt113    }114    115    if dpm_scheduler:116        scheduler_key = "dpm_scheduler"117    else:118        scheduler_key = "scheduler"119        120    image_np, so_img_list, b64 = run_ours(121        spec, bg_seed=seed, overall_prompt_override=overall_prompt_override, fg_seed_start=fg_seed_start, 122        fg_blending_ratio=fg_blending_ratio,frozen_step_ratio=frozen_step_ratio, use_autocast=use_autocast,123        gligen_scheduled_sampling_beta=gligen_scheduled_sampling_beta, num_inference_steps=num_inference_steps, scheduler_key=scheduler_key,124        so_negative_prompt=so_negative_prompt, overall_negative_prompt=overall_negative_prompt, so_batch_size=2125    )126    127    images = [image_np, b64]128    # if show_so_imgs:129    #     images.extend([np.asarray(so_img) for so_img in so_img_list])130    return images131 132def get_baseline_image(prompt, seed=0):133    if prompt == "":134        prompt = prompt_placeholder135    136    scheduler_key = "dpm_scheduler"137    num_inference_steps = 20138    139    image_np, b64 = run_baseline(prompt, bg_seed=seed, scheduler_key=scheduler_key, num_inference_steps=num_inference_steps)140    images = [image_np, b64]141    return images142 143def parse_input(text=None):144    try:145        if "Objects: " in text:146            text = text.split("Objects: ")[1]147            148        text_split = text.split(bg_prompt_text)149        if len(text_split) == 2:150            gen_boxes, bg_prompt = text_split151        gen_boxes = ast.literal_eval(gen_boxes.strip())    152        bg_prompt = bg_prompt.strip()153    except Exception as e:154        raise gr.Error(f"response format invalid: {e} (text: {text})")155    156    return gen_boxes, bg_prompt157 158def draw_boxes(anns):159    ax = plt.gca()160    ax.set_autoscale_on(False)161    polygons = []162    color = []163    for ann in anns:164        c = (np.random.random((1, 3))*0.6+0.4)165        [bbox_x, bbox_y, bbox_w, bbox_h] = ann['bbox']166        poly = [[bbox_x, bbox_y], [bbox_x, bbox_y+bbox_h],167                [bbox_x+bbox_w, bbox_y+bbox_h], [bbox_x+bbox_w, bbox_y]]168        np_poly = np.array(poly).reshape((4, 2))169        polygons.append(Polygon(np_poly))170        color.append(c)171 172        # print(ann)173        name = ann['name'] if 'name' in ann else str(ann['category_id'])174        ax.text(bbox_x, bbox_y, name, style='italic',175                bbox={'facecolor': 'white', 'alpha': 0.7, 'pad': 5})176 177    p = PatchCollection(polygons, facecolor='none',178                        edgecolors=color, linewidths=2)179    ax.add_collection(p)180 181 182def show_boxes(gen_boxes, bg_prompt=None):183    anns = [{'name': gen_box[0], 'bbox': gen_box[1]}184            for gen_box in gen_boxes]185 186    # White background (to allow line to show on the edge)187    I = np.ones((size[0]+4, size[1]+4, 3), dtype=np.uint8) * 255188 189    plt.imshow(I)190    plt.axis('off')191 192    if bg_prompt is not None:193        ax = plt.gca()194        ax.text(0, 0, bg_prompt, style='italic',195                bbox={'facecolor': 'white', 'alpha': 0.7, 'pad': 5})196 197        c = np.zeros((1, 3))198        [bbox_x, bbox_y, bbox_w, bbox_h] = (0, 0, size[1], size[0])199        poly = [[bbox_x, bbox_y], [bbox_x, bbox_y+bbox_h],200                [bbox_x+bbox_w, bbox_y+bbox_h], [bbox_x+bbox_w, bbox_y]]201        np_poly = np.array(poly).reshape((4, 2))202        polygons = [Polygon(np_poly)]203        color = [c]204        p = PatchCollection(polygons, facecolor='none',205                            edgecolors=color, linewidths=2)206        ax.add_collection(p)207 208    draw_boxes(anns)209 210duplicate_html = '<a style="display:inline-block" href="https://huggingface.co/spaces/longlian/llm-grounded-diffusion?duplicate=true"><img src="https://img.shields.io/badge/-Duplicate%20Space-blue?labelColor=white&style=flat&logo=data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAABAAAAAQCAYAAAAf8/9hAAAAAXNSR0IArs4c6QAAAP5JREFUOE+lk7FqAkEURY+ltunEgFXS2sZGIbXfEPdLlnxJyDdYB62sbbUKpLbVNhyYFzbrrA74YJlh9r079973psed0cvUD4A+4HoCjsA85X0Dfn/RBLBgBDxnQPfAEJgBY+A9gALA4tcbamSzS4xq4FOQAJgCDwV2CPKV8tZAJcAjMMkUe1vX+U+SMhfAJEHasQIWmXNN3abzDwHUrgcRGmYcgKe0bxrblHEB4E/pndMazNpSZGcsZdBlYJcEL9Afo75molJyM2FxmPgmgPqlWNLGfwZGG6UiyEvLzHYDmoPkDDiNm9JR9uboiONcBXrpY1qmgs21x1QwyZcpvxt9NS09PlsPAAAAAElFTkSuQmCC&logoWidth=14" alt="Duplicate Space"></a>'211 212html = f"""<h1>LLM-grounded Diffusion: Enhancing Prompt Understanding of Text-to-Image Diffusion Models with Large Language Models</h1>213            <h2>LLM + Stable Diffusion => better prompt understanding in text2image generation 🤩</h2>214            <h2><a href='https://llm-grounded-diffusion.github.io/'>Project Page</a> | <a href='https://bair.berkeley.edu/blog/2023/05/23/lmd/'>5-minute Blog Post</a> | <a href='https://arxiv.org/pdf/2305.13655.pdf'>ArXiv Paper</a> | <a href='https://github.com/TonyLianLong/LLM-groundedDiffusion'>Github</a> | <a href='https://llm-grounded-diffusion.github.io/#citation'>Cite our work</a> if our ideas inspire you.</h2>215            <p><b>Tips:</b><p>216            <p>1. If ChatGPT doesn't generate layout, add/remove the trailing space (added by default) and/or use GPT-4.</p>217            <p>2. You can perform multi-round specification by giving ChatGPT follow-up requests (e.g., make the object boxes bigger).</p>218            <p>3. You can also try prompts in Simplified Chinese. If you want to try prompts in another language, translate the first line of last example to your language.</p>219            <p>4. The diffusion model only runs 20 steps by default in this demo. You can make it run more steps to get higher quality images (or tweak frozen steps/guidance steps for better guidance and coherence).</p>220            <p>5. Duplicate this space and add GPU or clone the space and run locally to skip the queue and run our model faster. (<b>Currently we are using a T4 GPU on this space, which is quite slow, and you can add a A10G to make it 5x faster</b>) {duplicate_html}</p>221            <br/>222            <p>Implementation note: In this demo, we replace the attention manipulation in our layout-guided Stable Diffusion described in our paper with GLIGEN due to much faster inference speed (<b>FlashAttention supported, no backprop needed</b> during inference). Compared to vanilla GLIGEN, we have better coherence. Other parts of text-to-image pipeline, including single object generation and SAM, remain the same. The settings and examples in the prompt are simplified in this demo.</p>223            <style>.btn {{flex-grow: unset !important;}} </style>224            """225 226with gr.Blocks(227    title="LLM-grounded Diffusion: Enhancing Prompt Understanding of Text-to-Image Diffusion Models with Large Language Models"228) as g:229    gr.HTML(html)230    with gr.Tab("Stage 1. Image Prompt to ChatGPT"):231        with gr.Row():232            with gr.Column(scale=1):233                prompt = gr.Textbox(lines=2, label="Prompt for Layout Generation", placeholder=prompt_placeholder)234                generate_btn = gr.Button("Generate Prompt", variant='primary', elem_classes="btn")235                with gr.Accordion("Advanced options", open=False):236                    template = gr.Textbox(lines=10, label="Custom Template", placeholder="Customized Template", value=default_template)237            with gr.Column(scale=1):238                output = gr.Textbox(label="Paste this into ChatGPT (GPT-4 preferred; on Mac, click text and press Command+A and Command+C to copy all)", show_copy_button=True)239                gr.HTML("<a href='https://chat.openai.com' target='_blank'>Click here to open ChatGPT</a>")240        generate_btn.click(fn=get_lmd_prompt, inputs=[prompt, template], outputs=output, api_name="get_lmd_prompt")241    242        gr.Examples(243            examples=stage1_examples,244            inputs=[prompt],245            outputs=[output],246            fn=get_lmd_prompt,247            # cache_examples=True248        )249    250    with gr.Tab("Stage 2 (New). Layout to Image generation"):251        with gr.Row():252            with gr.Column(scale=1):253                response = gr.Textbox(lines=8, label="Paste ChatGPT response here (no original caption needed)", placeholder=layout_placeholder)254                overall_prompt_override = gr.Textbox(lines=2, label="Prompt for overall generation (optional but recommended)", placeholder="You can put your input prompt for layout generation here, helpful if your scene cannot be represented by background prompt and boxes only, e.g., with object interactions. If left empty: background prompt with [objects].", value="")255                num_inference_steps = gr.Slider(1, 250, value=20, step=1, label="Number of denoising steps (set to >=50 for higher generation quality)")256                seed = gr.Slider(0, 10000, value=0, step=1, label="Seed")257                with gr.Accordion("Advanced options (play around for better generation)", open=False):258                    frozen_step_ratio = gr.Slider(0, 1, value=0.4, step=0.1, label="Foreground frozen steps ratio (higher: preserve object attributes; lower: higher coherence; set to 0: (almost) equivalent to vanilla GLIGEN except details)")259                    gligen_scheduled_sampling_beta = gr.Slider(0, 1, value=0.3, step=0.1, label="GLIGEN guidance steps ratio (the beta value)")260                    dpm_scheduler = gr.Checkbox(label="Use DPM scheduler (unchecked: DDIM scheduler, may have better coherence, recommend >=50 inference steps)", show_label=False, value=True)261                    use_autocast = gr.Checkbox(label="Use FP16 Mixed Precision (faster but with slightly lower quality)", show_label=False, value=True)262                    fg_seed_start = gr.Slider(0, 10000, value=20, step=1, label="Seed for foreground variation")263                    fg_blending_ratio = gr.Slider(0, 1, value=0.1, step=0.01, label="Variations added to foreground for single object generation (0: no variation, 1: max variation)")264                    so_negative_prompt = gr.Textbox(lines=1, label="Negative prompt for single object generation", value=DEFAULT_SO_NEGATIVE_PROMPT)265                    overall_negative_prompt = gr.Textbox(lines=1, label="Negative prompt for overall generation", value=DEFAULT_OVERALL_NEGATIVE_PROMPT)266                    show_so_imgs = gr.Checkbox(label="Show annotated single object generations", show_label=False, value=False)267                    scale_boxes = gr.Checkbox(label="Scale bounding boxes to just fit the scene", show_label=False, value=False)268                visualize_btn = gr.Button("Visualize Layout", elem_classes="btn")269                generate_btn = gr.Button("Generate Image from Layout", variant='primary', elem_classes="btn")270            with gr.Column(scale=1):271                gallery = gr.Image(272                    label="Generated image", show_label=False, elem_id="gallery", columns=[1], rows=[1], object_fit="contain"                    273                )274                b64 = gr.Textbox(label="base64", placeholder="base64", lines = 2)275        visualize_btn.click(fn=get_layout_image_gallery, inputs=response, outputs=[gallery, b64], api_name="visualize-layout")276        generate_btn.click(fn=get_ours_image, inputs=[response, overall_prompt_override, seed, num_inference_steps, dpm_scheduler, use_autocast, fg_seed_start, fg_blending_ratio, frozen_step_ratio, gligen_scheduled_sampling_beta, so_negative_prompt, overall_negative_prompt, show_so_imgs, scale_boxes], outputs=[gallery, b64], api_name="layout-to-image")277 278        gr.Examples(279            examples=stage2_examples,280            inputs=[response, overall_prompt_override, seed],281            outputs=[gallery],282            fn=get_ours_image,283            # cache_examples=True284        )285 286    with gr.Tab("Baseline: Stable Diffusion"):287        with gr.Row():288            with gr.Column(scale=1):289                sd_prompt = gr.Textbox(lines=2, label="Prompt for baseline SD", placeholder=prompt_placeholder)290                seed = gr.Slider(0, 10000, value=0, step=1, label="Seed")291                generate_btn = gr.Button("Generate", elem_classes="btn")292            # with gr.Column(scale=1):293            #     output = gr.Image(shape=(512, 512), elem_classes="img", elem_id="img")294            with gr.Column(scale=1):295                gallery = gr.Image(296                    label="Generated image", show_label=False, elem_id="gallery", columns=[1], rows=[1], object_fit="contain"                    297                )298                b64 = gr.Textbox(label="base64", placeholder="base64", lines = 2)299        generate_btn.click(fn=get_baseline_image, inputs=[sd_prompt, seed], outputs=[gallery,b64], api_name="baseline")300 301        gr.Examples(302            examples=stage1_examples,303            inputs=[sd_prompt],304            outputs=[gallery],305            fn=get_baseline_image,306            # cache_examples=True307        )308 309g.launch()310