Team Ai
Modelpublic

ragavsachdeva/magi

sourceHugging Faceupdated 2y agoView on Hugging Face
58likes1.9kdownloads
utils.py395 linesDownload Raw Back to root
1import torch2import numpy as np3import random4import matplotlib.pyplot as plt5import matplotlib.patches as patches6from shapely.geometry import Point, box7import networkx as nx8from copy import deepcopy9from itertools import groupby10 11def move_to_device(inputs, device):12    if hasattr(inputs, "keys"):13        return {k: move_to_device(v, device) for k, v in inputs.items()}14    elif isinstance(inputs, list):15        return [move_to_device(v, device) for v in inputs]16    elif isinstance(inputs, tuple):17        return tuple([move_to_device(v, device) for v in inputs])18    elif isinstance(inputs, np.ndarray):19        return torch.from_numpy(inputs).to(device)20    else:21        return inputs.to(device)22 23class UnionFind:24    def __init__(self, n):25        self.parent = list(range(n))26        self.size = [1] * n27        self.num_components = n28 29    @classmethod30    def from_adj_matrix(cls, adj_matrix):31        ufds = cls(adj_matrix.shape[0])32        for i in range(adj_matrix.shape[0]):33            for j in range(adj_matrix.shape[1]):34                if adj_matrix[i, j] > 0:35                    ufds.unite(i, j)36        return ufds37    38    @classmethod39    def from_adj_list(cls, adj_list):40        ufds = cls(len(adj_list))41        for i in range(len(adj_list)):42            for j in adj_list[i]:43                ufds.unite(i, j)44        return ufds45    46    @classmethod47    def from_edge_list(cls, edge_list, num_nodes):48        ufds = cls(num_nodes)49        for edge in edge_list:50            ufds.unite(edge[0], edge[1])51        return ufds52 53    def find(self, x):54        if self.parent[x] == x:55            return x56        self.parent[x] = self.find(self.parent[x])57        return self.parent[x]58 59    def unite(self, x, y):60        x = self.find(x)61        y = self.find(y)62        if x != y:63            if self.size[x] < self.size[y]:64                x, y = y, x65            self.parent[y] = x66            self.size[x] += self.size[y]67            self.num_components -= 168    69    def get_components_of(self, x):70        x = self.find(x)71        return [i for i in range(len(self.parent)) if self.find(i) == x]72    73    def are_connected(self, x, y):74        return self.find(x) == self.find(y)75 76    def get_size(self, x):77        return self.size[self.find(x)]78 79    def get_num_components(self):80        return self.num_components81    82    def get_labels_for_connected_components(self):83        map_parent_to_label = {}84        labels = []85        for i in range(len(self.parent)):86            parent = self.find(i)87            if parent not in map_parent_to_label:88                map_parent_to_label[parent] = len(map_parent_to_label)89            labels.append(map_parent_to_label[parent])90        return labels91 92def visualise_single_image_prediction(image_as_np_array, predictions, filename):93    h, w = image_as_np_array.shape[:2]94    if h > w:95        figure, subplot = plt.subplots(1, 1, figsize=(10, 10 * h / w))96    else:97        figure, subplot = plt.subplots(1, 1, figsize=(10 * w / h, 10))98    subplot.imshow(image_as_np_array)99    plot_bboxes(subplot, predictions["panels"], color="green")100    plot_bboxes(subplot, predictions["texts"], color="red", add_index=True)101    plot_bboxes(subplot, predictions["characters"], color="blue")102 103    COLOURS = [104        "#b7ff51", # green105        "#f50a8f", # pink106        "#4b13b6", # purple107        "#ddaa34", # orange108        "#bea2a2", # brown109    ]110    colour_index = 0111    character_cluster_labels = predictions["character_cluster_labels"]112    unique_label_sorted_by_frequency = sorted(list(set(character_cluster_labels)), key=lambda x: character_cluster_labels.count(x), reverse=True)113    for label in unique_label_sorted_by_frequency:114        root = None115        others = []116        for i in range(len(predictions["characters"])):117            if character_cluster_labels[i] == label:118                if root is None:119                    root = i120                else:121                    others.append(i)122        if colour_index >= len(COLOURS):123            random_colour = COLOURS[0]124            while random_colour in COLOURS:125                random_colour = "#" + "".join([random.choice("0123456789ABCDEF") for j in range(6)])126        else:127            random_colour = COLOURS[colour_index]128            colour_index += 1129        bbox_i = predictions["characters"][root]130        x1 = bbox_i[0] + (bbox_i[2] - bbox_i[0]) / 2131        y1 = bbox_i[1] + (bbox_i[3] - bbox_i[1]) / 2132        subplot.plot([x1], [y1], color=random_colour, marker="o", markersize=5)133        for j in others:134            # draw line from centre of bbox i to centre of bbox j135            bbox_j = predictions["characters"][j]136            x1 = bbox_i[0] + (bbox_i[2] - bbox_i[0]) / 2137            y1 = bbox_i[1] + (bbox_i[3] - bbox_i[1]) / 2138            x2 = bbox_j[0] + (bbox_j[2] - bbox_j[0]) / 2139            y2 = bbox_j[1] + (bbox_j[3] - bbox_j[1]) / 2140            subplot.plot([x1, x2], [y1, y2], color=random_colour, linewidth=2)141            subplot.plot([x2], [y2], color=random_colour, marker="o", markersize=5)142    143    for (i, j) in predictions["text_character_associations"]:144        score = predictions["dialog_confidences"][i]145        bbox_i = predictions["texts"][i]146        bbox_j = predictions["characters"][j]147        x1 = bbox_i[0] + (bbox_i[2] - bbox_i[0]) / 2148        y1 = bbox_i[1] + (bbox_i[3] - bbox_i[1]) / 2149        x2 = bbox_j[0] + (bbox_j[2] - bbox_j[0]) / 2150        y2 = bbox_j[1] + (bbox_j[3] - bbox_j[1]) / 2151        subplot.plot([x1, x2], [y1, y2], color="red", linewidth=2, linestyle="dashed", alpha=score)152 153    subplot.axis("off")154    if filename is not None:155        plt.savefig(filename, bbox_inches="tight", pad_inches=0)156 157    figure.canvas.draw()158    image = np.array(figure.canvas.renderer._renderer)159    plt.close()160    return image161 162def plot_bboxes(subplot, bboxes, color="red", add_index=False):163    for id, bbox in enumerate(bboxes):164        w = bbox[2] - bbox[0]165        h = bbox[3] - bbox[1]166        rect = patches.Rectangle(167            bbox[:2], w, h, linewidth=1, edgecolor=color, facecolor="none", linestyle="solid"168        )169        subplot.add_patch(rect)170        if add_index:171            cx, cy = bbox[0] + w / 2, bbox[1] + h / 2172            subplot.text(cx, cy, str(id), color=color, fontsize=10, ha="center", va="center")173 174def sort_panels(rects):175    before_rects = convert_to_list_of_lists(rects)176    # slightly erode all rectangles initially to account for imperfect detections177    rects = [erode_rectangle(rect, 0.05) for rect in before_rects]178    G = nx.DiGraph()179    G.add_nodes_from(range(len(rects)))180    for i in range(len(rects)):181        for j in range(len(rects)):182            if i == j:183                continue184            if is_there_a_directed_edge(i, j, rects):185                G.add_edge(i, j, weight=get_distance(rects[i], rects[j]))186            else:187                G.add_edge(j, i, weight=get_distance(rects[i], rects[j]))188    while True:189        cycles = sorted(nx.simple_cycles(G))190        cycles = [cycle for cycle in cycles if len(cycle) > 1]191        if len(cycles) == 0:192            break193        cycle = cycles[0]194        edges = [e for e in zip(cycle, cycle[1:] + cycle[:1])]195        max_cyclic_edge = max(edges, key=lambda x: G.edges[x]["weight"])196        G.remove_edge(*max_cyclic_edge)197    return list(nx.topological_sort(G))198 199def is_strictly_above(rectA, rectB):200    x1A, y1A, x2A, y2A = rectA201    x1B, y1B, x2B, y2B = rectB202    return y2A < y1B203 204def is_strictly_below(rectA, rectB):205    x1A, y1A, x2A, y2A = rectA206    x1B, y1B, x2B, y2B = rectB207    return y2B < y1A208 209def is_strictly_left_of(rectA, rectB):210    x1A, y1A, x2A, y2A = rectA211    x1B, y1B, x2B, y2B = rectB212    return x2A < x1B213 214def is_strictly_right_of(rectA, rectB):215    x1A, y1A, x2A, y2A = rectA216    x1B, y1B, x2B, y2B = rectB217    return x2B < x1A218 219def intersects(rectA, rectB):220    return box(*rectA).intersects(box(*rectB))221 222def is_there_a_directed_edge(a, b, rects):223    rectA = rects[a]224    rectB = rects[b]225    centre_of_A = [rectA[0] + (rectA[2] - rectA[0]) / 2, rectA[1] + (rectA[3] - rectA[1]) / 2]226    centre_of_B = [rectB[0] + (rectB[2] - rectB[0]) / 2, rectB[1] + (rectB[3] - rectB[1]) / 2]227    if np.allclose(np.array(centre_of_A), np.array(centre_of_B)):228        return box(*rectA).area > (box(*rectB)).area229    copy_A = [rectA[0], rectA[1], rectA[2], rectA[3]]230    copy_B = [rectB[0], rectB[1], rectB[2], rectB[3]]231    while True:232        if is_strictly_above(copy_A, copy_B) and not is_strictly_left_of(copy_A, copy_B):233            return 1234        if is_strictly_above(copy_B, copy_A) and not is_strictly_left_of(copy_B, copy_A):235            return 0236        if is_strictly_right_of(copy_A, copy_B) and not is_strictly_below(copy_A, copy_B):237            return 1238        if is_strictly_right_of(copy_B, copy_A) and not is_strictly_below(copy_B, copy_A):239            return 0240        if is_strictly_below(copy_A, copy_B) and is_strictly_right_of(copy_A, copy_B):241            return use_cuts_to_determine_edge_from_a_to_b(a, b, rects)242        if is_strictly_below(copy_B, copy_A) and is_strictly_right_of(copy_B, copy_A):243           return use_cuts_to_determine_edge_from_a_to_b(a, b, rects)244        # otherwise they intersect245        copy_A = erode_rectangle(copy_A, 0.05)246        copy_B = erode_rectangle(copy_B, 0.05)247    248def get_distance(rectA, rectB):249    return box(rectA[0], rectA[1], rectA[2], rectA[3]).distance(box(rectB[0], rectB[1], rectB[2], rectB[3]))250 251def use_cuts_to_determine_edge_from_a_to_b(a, b, rects):252    rects = deepcopy(rects)253    while True:254        xmin, ymin, xmax, ymax = min(rects[a][0], rects[b][0]), min(rects[a][1], rects[b][1]), max(rects[a][2], rects[b][2]), max(rects[a][3], rects[b][3])255        rect_index = [i for i in range(len(rects)) if intersects(rects[i], [xmin, ymin, xmax, ymax])]256        rects_copy = [rect for rect in rects if intersects(rect, [xmin, ymin, xmax, ymax])]257        258        # try to split the panels using a "horizontal" lines259        overlapping_y_ranges = merge_overlapping_ranges([(y1, y2) for x1, y1, x2, y2 in rects_copy])260        panel_index_to_split = {}261        for split_index, (y1, y2) in enumerate(overlapping_y_ranges):262            for i, index in enumerate(rect_index):263                if y1 <= rects_copy[i][1] <= rects_copy[i][3] <= y2:264                    panel_index_to_split[index] = split_index265        266        if panel_index_to_split[a] != panel_index_to_split[b]:267            return panel_index_to_split[a] < panel_index_to_split[b]268        269        # try to split the panels using a "vertical" lines270        overlapping_x_ranges = merge_overlapping_ranges([(x1, x2) for x1, y1, x2, y2 in rects_copy])271        panel_index_to_split = {}272        for split_index, (x1, x2) in enumerate(overlapping_x_ranges[::-1]):273            for i, index in enumerate(rect_index):274                if x1 <= rects_copy[i][0] <= rects_copy[i][2] <= x2:275                    panel_index_to_split[index] = split_index276        if panel_index_to_split[a] != panel_index_to_split[b]:277            return panel_index_to_split[a] < panel_index_to_split[b]278        279        # otherwise, erode the rectangles and try again280        rects = [erode_rectangle(rect, 0.05) for rect in rects]281 282def erode_rectangle(bbox, erosion_factor):283    x1, y1, x2, y2 = bbox284    w, h = x2 - x1, y2 - y1285    cx, cy = x1 + w / 2, y1 + h / 2286    if w < h:287        aspect_ratio = w / h288        erosion_factor_width = erosion_factor * aspect_ratio289        erosion_factor_height = erosion_factor290    else:291        aspect_ratio = h / w292        erosion_factor_width = erosion_factor293        erosion_factor_height = erosion_factor * aspect_ratio294    w = w - w * erosion_factor_width295    h = h - h * erosion_factor_height296    x1, y1, x2, y2 = cx - w / 2, cy - h / 2, cx + w / 2, cy + h / 2297    return [x1, y1, x2, y2]298 299def merge_overlapping_ranges(ranges):300    """301    ranges: list of tuples (x1, x2)302    """303    if len(ranges) == 0:304        return []305    ranges = sorted(ranges, key=lambda x: x[0])306    merged_ranges = []307    for i, r in enumerate(ranges):308        if i == 0:309            prev_x1, prev_x2 = r310            continue311        x1, x2 = r312        if x1 > prev_x2:313            merged_ranges.append((prev_x1, prev_x2))314            prev_x1, prev_x2 = x1, x2315        else:316            prev_x2 = max(prev_x2, x2)317    merged_ranges.append((prev_x1, prev_x2))318    return merged_ranges319 320def sort_text_boxes_in_reading_order(text_bboxes, sorted_panel_bboxes):321    text_bboxes = convert_to_list_of_lists(text_bboxes)322    sorted_panel_bboxes = convert_to_list_of_lists(sorted_panel_bboxes)323 324    if len(text_bboxes) == 0:325        return []326 327    def indices_of_same_elements(nums):328        groups = groupby(range(len(nums)), key=lambda i: nums[i])329        return [list(indices) for _, indices in groups]330 331    panel_id_for_text = get_text_to_panel_mapping(text_bboxes, sorted_panel_bboxes)332    indices_of_texts = list(range(len(text_bboxes)))333    indices_of_texts, panel_id_for_text = zip(*sorted(zip(indices_of_texts, panel_id_for_text), key=lambda x: x[1]))334    indices_of_texts = list(indices_of_texts)335    grouped_indices = indices_of_same_elements(panel_id_for_text)336    for group in grouped_indices:337        subset_of_text_indices = [indices_of_texts[i] for i in group]338        text_bboxes_of_subset = [text_bboxes[i] for i in subset_of_text_indices]339        sorted_subset_indices = sort_texts_within_panel(text_bboxes_of_subset)340        indices_of_texts[group[0] : group[-1] + 1] = [subset_of_text_indices[i] for i in sorted_subset_indices]341    return indices_of_texts342 343def get_text_to_panel_mapping(text_bboxes, sorted_panel_bboxes):344    text_to_panel_mapping = []345    for text_bbox in text_bboxes:346        shapely_text_polygon = box(*text_bbox)347        all_intersections = []348        all_distances = []349        if len(sorted_panel_bboxes) == 0:350            text_to_panel_mapping.append(-1)351            continue352        for j, annotation in enumerate(sorted_panel_bboxes):353            shapely_annotation_polygon = box(*annotation)354            if shapely_text_polygon.intersects(shapely_annotation_polygon):355                all_intersections.append((shapely_text_polygon.intersection(shapely_annotation_polygon).area, j))356            all_distances.append((shapely_text_polygon.distance(shapely_annotation_polygon), j))357        if len(all_intersections) == 0:358            text_to_panel_mapping.append(min(all_distances, key=lambda x: x[0])[1])359        else:360            text_to_panel_mapping.append(max(all_intersections, key=lambda x: x[0])[1])361    return text_to_panel_mapping362 363def sort_texts_within_panel(rects):364    smallest_y = float("inf")365    greatest_x = float("-inf")366    for i, rect in enumerate(rects):367        x1, y1, x2, y2 = rect368        smallest_y = min(smallest_y, y1)369        greatest_x = max(greatest_x, x2)370    371    reference_point = Point(greatest_x, smallest_y)372 373    polygons_and_index = []374    for i, rect in enumerate(rects):375        x1, y1, x2, y2 = rect376        polygons_and_index.append((box(x1,y1,x2,y2), i))377    # sort points by closest to reference point378    polygons_and_index = sorted(polygons_and_index, key=lambda x: reference_point.distance(x[0]))379    indices = [x[1] for x in polygons_and_index]380    return indices381 382def x1y1wh_to_x1y1x2y2(bbox):383    x1, y1, w, h = bbox384    return [x1, y1, x1 + w, y1 + h]385 386def x1y1x2y2_to_xywh(bbox):387    x1, y1, x2, y2 = bbox388    return [x1, y1, x2 - x1, y2 - y1]389 390def convert_to_list_of_lists(rects):391    if isinstance(rects, torch.Tensor):392        return rects.tolist()393    if isinstance(rects, np.ndarray):394        return rects.tolist()395    return [[a, b, c, d] for a, b, c, d in rects]