ragavsachdeva/magi
581.9k
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]