AnchoredAI/llm-grounded-diffusion
0
1import ast2import os3import json4from matplotlib.patches import Polygon5from matplotlib.collections import PatchCollection6import matplotlib.pyplot as plt7import numpy as np8import cv29import inflect10 11p = inflect.engine()12 13img_dir = "imgs"14bg_prompt_text = "Background prompt: "15# h, w16box_scale = (512, 512)17size = box_scale18size_h, size_w = size19print(f"Using box scale: {box_scale}")20 21def parse_input(text=None, no_input=False):22 if not text:23 if no_input:24 return25 26 text = input("Enter the response: ")27 if "Objects: " in text:28 text = text.split("Objects: ")[1]29 30 text_split = text.split(bg_prompt_text)31 if len(text_split) == 2:32 gen_boxes, bg_prompt = text_split33 elif len(text_split) == 1:34 if no_input:35 return36 gen_boxes = text37 bg_prompt = ""38 while not bg_prompt:39 # Ignore the empty lines in the response40 bg_prompt = input("Enter the background prompt: ").strip()41 if bg_prompt_text in bg_prompt:42 bg_prompt = bg_prompt.split(bg_prompt_text)[1]43 else:44 raise ValueError(f"text: {text}")45 try:46 gen_boxes = ast.literal_eval(gen_boxes) 47 except SyntaxError as e:48 # Sometimes the response is in plain text49 if "No objects" in gen_boxes:50 gen_boxes = []51 else:52 raise e53 bg_prompt = bg_prompt.strip()54 55 return gen_boxes, bg_prompt56 57def filter_boxes(gen_boxes, scale_boxes=True, ignore_background=True, max_scale=3):58 if len(gen_boxes) == 0:59 return []60 61 box_dict_format = False62 gen_boxes_new = []63 for gen_box in gen_boxes:64 if isinstance(gen_box, dict):65 name, [bbox_x, bbox_y, bbox_w, bbox_h] = gen_box['name'], gen_box['bounding_box']66 box_dict_format = True67 else:68 name, [bbox_x, bbox_y, bbox_w, bbox_h] = gen_box69 if bbox_w <= 0 or bbox_h <= 0:70 # Empty boxes71 continue72 if ignore_background:73 if (bbox_w >= size[1] and bbox_h >= size[0]) or bbox_x > size[1] or bbox_y > size[0]:74 # Ignore the background boxes75 continue76 gen_boxes_new.append(gen_box)77 78 gen_boxes = gen_boxes_new79 80 if len(gen_boxes) == 0:81 return []82 83 filtered_gen_boxes = []84 if box_dict_format:85 # For compatibility86 bbox_left_x_min = min([gen_box['bounding_box'][0] for gen_box in gen_boxes])87 bbox_right_x_max = max([gen_box['bounding_box'][0] + gen_box['bounding_box'][2] for gen_box in gen_boxes])88 bbox_top_y_min = min([gen_box['bounding_box'][1] for gen_box in gen_boxes])89 bbox_bottom_y_max = max([gen_box['bounding_box'][1] + gen_box['bounding_box'][3] for gen_box in gen_boxes])90 else:91 bbox_left_x_min = min([gen_box[1][0] for gen_box in gen_boxes])92 bbox_right_x_max = max([gen_box[1][0] + gen_box[1][2] for gen_box in gen_boxes])93 bbox_top_y_min = min([gen_box[1][1] for gen_box in gen_boxes])94 bbox_bottom_y_max = max([gen_box[1][1] + gen_box[1][3] for gen_box in gen_boxes])95 96 # All boxes are empty97 if (bbox_right_x_max - bbox_left_x_min) == 0:98 return []99 100 # Used if scale_boxes is True101 shift = -bbox_left_x_min102 scale = size_w / (bbox_right_x_max - bbox_left_x_min)103 104 scale = min(scale, max_scale)105 106 for gen_box in gen_boxes:107 if box_dict_format:108 name, [bbox_x, bbox_y, bbox_w, bbox_h] = gen_box['name'], gen_box['bounding_box']109 else:110 name, [bbox_x, bbox_y, bbox_w, bbox_h] = gen_box111 112 if scale_boxes:113 # Vertical: move the boxes if out of bound114 # Horizontal: move and scale the boxes so it spans the horizontal line115 116 bbox_x = (bbox_x + shift) * scale117 bbox_y = bbox_y * scale118 bbox_w, bbox_h = bbox_w * scale, bbox_h * scale119 # TODO: verify this makes the y center not moving120 bbox_y_offset = 0121 if bbox_top_y_min * scale + bbox_y_offset < 0:122 bbox_y_offset -= bbox_top_y_min * scale123 if bbox_bottom_y_max * scale + bbox_y_offset >= size_h:124 bbox_y_offset -= bbox_bottom_y_max * scale - size_h125 bbox_y += bbox_y_offset126 127 if bbox_y < 0:128 bbox_y, bbox_h = 0, bbox_h - bbox_y129 130 name = name.rstrip(".")131 bounding_box = (int(np.round(bbox_x)), int(np.round(bbox_y)), int(np.round(bbox_w)), int(np.round(bbox_h)))132 if box_dict_format:133 gen_box = {134 'name': name,135 'bounding_box': bounding_box136 }137 else:138 gen_box = (name, bounding_box)139 140 filtered_gen_boxes.append(gen_box)141 142 return filtered_gen_boxes143 144def draw_boxes(anns):145 ax = plt.gca()146 ax.set_autoscale_on(False)147 polygons = []148 color = []149 for ann in anns:150 c = (np.random.random((1, 3))*0.6+0.4)151 [bbox_x, bbox_y, bbox_w, bbox_h] = ann['bbox']152 poly = [[bbox_x, bbox_y], [bbox_x, bbox_y+bbox_h],153 [bbox_x+bbox_w, bbox_y+bbox_h], [bbox_x+bbox_w, bbox_y]]154 np_poly = np.array(poly).reshape((4, 2))155 polygons.append(Polygon(np_poly))156 color.append(c)157 158 # print(ann)159 name = ann['name'] if 'name' in ann else str(ann['category_id'])160 ax.text(bbox_x, bbox_y, name, style='italic',161 bbox={'facecolor': 'white', 'alpha': 0.7, 'pad': 5})162 163 p = PatchCollection(polygons, facecolor='none',164 edgecolors=color, linewidths=2)165 ax.add_collection(p)166 167 168def show_boxes(gen_boxes, bg_prompt=None, ind=None, show=False):169 if len(gen_boxes) == 0:170 return171 172 if isinstance(gen_boxes[0], dict):173 anns = [{'name': gen_box['name'], 'bbox': gen_box['bounding_box']}174 for gen_box in gen_boxes]175 else:176 anns = [{'name': gen_box[0], 'bbox': gen_box[1]} for gen_box in gen_boxes]177 178 # White background (to allow line to show on the edge)179 I = np.ones((size[0]+4, size[1]+4, 3), dtype=np.uint8) * 255180 181 plt.imshow(I)182 plt.axis('off')183 184 if bg_prompt is not None:185 ax = plt.gca()186 ax.text(0, 0, bg_prompt, style='italic',187 bbox={'facecolor': 'white', 'alpha': 0.7, 'pad': 5})188 189 c = (np.zeros((1, 3)))190 [bbox_x, bbox_y, bbox_w, bbox_h] = (0, 0, size[1], size[0])191 poly = [[bbox_x, bbox_y], [bbox_x, bbox_y+bbox_h],192 [bbox_x+bbox_w, bbox_y+bbox_h], [bbox_x+bbox_w, bbox_y]]193 np_poly = np.array(poly).reshape((4, 2))194 polygons = [Polygon(np_poly)]195 color = [c]196 p = PatchCollection(polygons, facecolor='none',197 edgecolors=color, linewidths=2)198 ax.add_collection(p)199 200 draw_boxes(anns)201 if show:202 plt.show()203 else:204 print("Saved to", f"{img_dir}/boxes.png", f"ind: {ind}")205 if ind is not None:206 plt.savefig(f"{img_dir}/boxes_{ind}.png")207 plt.savefig(f"{img_dir}/boxes.png")208 209 210def show_masks(masks):211 masks_to_show = np.zeros((*size, 3), dtype=np.float32)212 for mask in masks:213 c = (np.random.random((3,))*0.6+0.4)214 215 masks_to_show += mask[..., None] * c[None, None, :]216 plt.imshow(masks_to_show)217 plt.savefig(f"{img_dir}/masks.png")218 plt.show()219 plt.clf()220 221def convert_box(box, height, width):222 # box: x, y, w, h (in 512 format) -> x_min, y_min, x_max, y_max223 x_min, y_min = box[0] / width, box[1] / height224 w_box, h_box = box[2] / width, box[3] / height225 226 x_max, y_max = x_min + w_box, y_min + h_box227 228 return x_min, y_min, x_max, y_max229 230def convert_spec(spec, height, width, include_counts=True, verbose=False):231 # Infer from spec232 prompt, gen_boxes, bg_prompt = spec['prompt'], spec['gen_boxes'], spec['bg_prompt']233 234 # This ensures the same objects appear together because flattened `overall_phrases_bboxes` should EXACTLY correspond to `so_prompt_phrase_box_list`. 235 gen_boxes = sorted(gen_boxes, key=lambda gen_box: gen_box[0])236 237 gen_boxes = [(name, convert_box(box, height=height, width=width)) for name, box in gen_boxes]238 239 # NOTE: so phrase should include all the words associated to the object (otherwise "an orange dog" may be recognized as "an orange" by the model generating the background).240 # so word should have one token that includes the word to transfer cross attention (the object name).241 # Currently using the last word of the object name as word.242 if bg_prompt:243 so_prompt_phrase_word_box_list = [(f"{bg_prompt} with {name}", name, name.split(" ")[-1], box) for name, box in gen_boxes]244 else:245 so_prompt_phrase_word_box_list = [(f"{name}", name, name.split(" ")[-1], box) for name, box in gen_boxes]246 247 objects = [gen_box[0] for gen_box in gen_boxes]248 249 objects_unique, objects_count = np.unique(objects, return_counts=True)250 251 num_total_matched_boxes = 0252 overall_phrases_words_bboxes = []253 for ind, object_name in enumerate(objects_unique):254 bboxes = [box for name, box in gen_boxes if name == object_name]255 256 if objects_count[ind] > 1:257 phrase = p.plural_noun(object_name.replace("an ", "").replace("a ", ""))258 if include_counts:259 phrase = p.number_to_words(objects_count[ind]) + " " + phrase260 else:261 phrase = object_name262 # Currently using the last word of the phrase as word.263 word = phrase.split(' ')[-1]264 265 num_total_matched_boxes += len(bboxes)266 overall_phrases_words_bboxes.append((phrase, word, bboxes))267 268 assert num_total_matched_boxes == len(gen_boxes), f"{num_total_matched_boxes} != {len(gen_boxes)}"269 270 objects_str = ", ".join([phrase for phrase, _, _ in overall_phrases_words_bboxes])271 if objects_str:272 if bg_prompt:273 overall_prompt = f"{bg_prompt} with {objects_str}"274 else:275 overall_prompt = objects_str276 else:277 overall_prompt = bg_prompt278 279 if verbose:280 print("so_prompt_phrase_word_box_list:", so_prompt_phrase_word_box_list)281 print("overall_prompt:", overall_prompt)282 print("overall_phrases_words_bboxes:", overall_phrases_words_bboxes)283 284 return so_prompt_phrase_word_box_list, overall_prompt, overall_phrases_words_bboxes285 