michaelcreatesstuff/llm-grounded-diffusion
2
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 