Team Ai
Apppublic

AnchoredAI/llm-grounded-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
parse.py285 linesDownload Raw Back to utils
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