import json from datetime import date import numpy as np from matplotlib import colormaps from PIL import ImageColor, ImageDraw, ImageFont today = date.today() FONTS = { "amiko": "fonts/Amiko-Regular.ttf", "nature": "fonts/LoveNature.otf", "painter": "fonts/PainterDecorator.otf", "animals": "fonts/UncialAnimals.ttf", "zen": "fonts/ZEN.TTF", } # perceptually uniform maps first; turbo separates neighbouring bodyparts best COLORMAPS = ["viridis", "plasma", "magma", "cividis", "turbo"] ######################################### # Draw keypoints on image def draw_keypoints_on_image( image, keypoints, map_label_id_to_str, flag_show_str_labels, use_normalized_coordinates=True, font_style="amiko", font_size=8, keypt_color="#ff0000", marker_size=2, color_by_confidence=True, colormap="viridis", ): """Draws keypoints on an image. Modified from: https://www.programcreek.com/python/?code=fjchange%2Fobject_centric_VAD%2Fobject_centric_VAD-master%2Fobject_detection%2Futils%2Fvisualization_utils.py Args: image: a PIL.Image object. keypoints: a numpy array with shape [num_keypoints, 2]. map_label_id_to_str: dict with keys=label number and values= label string flag_show_str_labels: boolean to select whether or not to show string labels color: color to draw the keypoints with. Default is red. radius: keypoint radius. Default value is 2. use_normalized_coordinates: if True (default), treat keypoint values as relative to the image. Otherwise treat them as absolute. """ # get a drawing context draw = ImageDraw.Draw(image, "RGBA") im_width, im_height = image.size keypoints_x = [k[0] for k in keypoints] keypoints_y = [k[1] for k in keypoints] confidences = [k[2] for k in keypoints] # adjust keypoints coords if required if use_normalized_coordinates: keypoints_x = tuple([im_width * x for x in keypoints_x]) keypoints_y = tuple([im_height * y for y in keypoints_y]) cmap = colormaps[colormap] # draw ellipses around keypoints for i, (keypoint_x, keypoint_y) in enumerate(zip(keypoints_x, keypoints_y, strict=True)): # handling potential nans in the keypoints if np.isnan(keypoint_x).any(): continue confidence = float(np.clip(confidences[i], 0, 1)) if color_by_confidence: # fill color encodes the keypoint confidence (see confidence_legend_html in ui_utils) round_fill = cmap(confidence, bytes=True) else: # one color per bodypart, transparency encodes the confidence round_fill = list(cmap(i / max(len(keypoints) - 1, 1), bytes=True)) round_fill[3] = round(confidence * 255) round_fill = tuple(round_fill) draw.ellipse( [ (keypoint_x - marker_size, keypoint_y - marker_size), (keypoint_x + marker_size, keypoint_y + marker_size), ], fill=tuple(round_fill), outline="black", width=1, ) # fill and outline: [0,255] # add string labels around keypoints if flag_show_str_labels: font = ImageFont.truetype(FONTS[font_style], font_size) draw.text( (keypoint_x + marker_size, keypoint_y + marker_size), # (0.5*im_width, 0.5*im_height), #------- display_bodypart(map_label_id_to_str[i]), ImageColor.getcolor(keypt_color, "RGB"), # rgb # font=font, ) ######################################### # Bodypart names for display # display names where the SuperAnimal definitions misspell (quadruped "thai") or read oddly # (top-view mouse "backend"); the JSON output keeps the model's names BODYPART_DISPLAY_NAMES = { "front_left_thai": "front left thigh", "front_right_thai": "front right thigh", "back_left_thai": "back left thigh", "back_right_thai": "back right thigh", "mid_backend": "mid back end", "mid_backend2": "mid back end 2", "mid_backend3": "mid back end 3", } def display_bodypart(name): return BODYPART_DISPLAY_NAMES.get(name, name.replace("_", " ")) ######################################### # Keypoint confidences as table rows def keypoint_confidence_rows(kpts_per_animal, map_label_id_to_str): """(animal, bodypart, confidence) for every keypoint kept (not NaN), lowest confidence first.""" rows = [] for i_animal, kpts in enumerate(kpts_per_animal): for i_kpt, kpt in enumerate(kpts): if not np.isnan(kpt[2]): rows.append([i_animal, display_bodypart(map_label_id_to_str[i_kpt]), round(float(kpt[2]), 3)]) return sorted(rows, key=lambda row: row[2]) ######################################### # Save the annotated image for download def save_annotated_image(image, path_to_output_file="download_annotated.png"): image.save(path_to_output_file) return path_to_output_file ######################################### # Draw bboxes on image def draw_bbox_w_text(img, results, font_style="amiko", font_size=8, bbox_color="#ff0000"): x1, y1, x2, y2, confidence = results[:5] draw = ImageDraw.Draw(img) draw.rectangle([(x1, y1), (x2, y2)], outline=bbox_color, width=max(2, round(font_size / 5))) label = f"animal {confidence:.2f}" font = ImageFont.truetype(FONTS[font_style], font_size) left, top, right, bottom = draw.textbbox((0, 0), label, font=font) pad = max(2, font_size // 5) label_w, label_h = right - left + 2 * pad, bottom - top + 2 * pad # label above the box, or inside it when the box touches the top of the image label_y = y1 - label_h if y1 >= label_h else y1 draw.rectangle([(x1, label_y), (x1 + label_w, label_y + label_h)], fill=bbox_color) draw.text((x1 + pad - left, label_y + pad - top), label, font=font, fill=label_text_color(bbox_color)) def label_text_color(background): # black or white, whichever contrasts more with the background (WCAG relative luminance) channels = [c / 255 for c in ImageColor.getrgb(background)[:3]] r, g, b = [c / 12.92 if c <= 0.03928 else ((c + 0.055) / 1.055) ** 2.4 for c in channels] return "black" if 0.2126 * r + 0.7152 * g + 0.0722 * b > 0.179 else "white" ########################################### def save_results_as_json( md_results, dlc_outputs, animal_bboxes, map_dlc_label_id_to_str, model, mega_model_input, path_to_output_file="download_predictions.json", ): """ Output detections as json file animal_bboxes: detection rows [x1,y1,x2,y2,conf,label], one per entry of dlc_outputs """ # initialise dict to save to json info = {} info["date"] = str(today) info["MD_model"] = str(mega_model_input) # info from megaDetector info["file"] = md_results.files[0] number_bb = len(md_results.xyxy[0].tolist()) info["number_of_bb"] = number_bb # info from DLC info["dlc_model"] = model labels = [n for n in map_dlc_label_id_to_str.values()] # define aux dict for every animal bounding box above threshold for i in range(len(dlc_outputs)): aux = {} # MD output corner_x1, corner_y1, corner_x2, corner_y2, confidence, _ = animal_bboxes[i] aux["corner_1"] = (corner_x1, corner_y1) aux["corner_2"] = (corner_x2, corner_y2) aux["predict MD"] = md_results.names[0] aux["confidence MD"] = confidence # DLC output kypts = [] for s in dlc_outputs[i]: aux1 = [] for j in s: aux1.append(float(j)) kypts.append(aux1) aux["dlc_pred"] = dict(zip(labels, kypts, strict=True)) info["bb_" + str(i)] = aux # save dict as json with open(path_to_output_file, "w") as f: json.dump(info, f, indent=1) print(f"Output file saved at {path_to_output_file}") return path_to_output_file def save_results_only_dlc(dlc_outputs, map_label_id_to_str, model, output_file="dowload_predictions_dlc.json"): """ write json dlc output """ info = {} info["date"] = str(today) labels = [n for n in map_label_id_to_str.values()] info["dlc_model"] = model kypts = [] for s in dlc_outputs: aux1 = [] for j in s: aux1.append(float(j)) kypts.append(aux1) info["dlc_pred"] = dict(zip(labels, kypts, strict=True)) with open(output_file, "w") as f: json.dump(info, f, indent=1) print(f"Output file saved at {output_file}") return output_file def save_results_pytorch( animals, map_label_id_to_str, model, pose_model, detector, path_to_output_file="download_predictions.json" ): """ Output PyTorch SuperAnimal predictions as json file (same layout as save_results_as_json) animals: list of {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3)}, in image coords detector: None if the detector was skipped (whole image used as one animal) """ info = {} info["date"] = str(today) info["backend"] = "pytorch" info["dlc_model"] = model info["pose_model"] = pose_model info["detector"] = detector info["number_of_bb"] = len(animals) labels = [n for n in map_label_id_to_str.values()] for i, animal in enumerate(animals): corner_x1, corner_y1, corner_x2, corner_y2, confidence = animal["bbox"] aux = {} aux["corner_1"] = (corner_x1, corner_y1) aux["corner_2"] = (corner_x2, corner_y2) aux["confidence"] = confidence aux["dlc_pred"] = dict(zip(labels, [[float(v) for v in kpt] for kpt in animal["kpts"]], strict=True)) info["bb_" + str(i)] = aux with open(path_to_output_file, "w") as f: json.dump(info, f, indent=1) print(f"Output file saved at {path_to_output_file}") return path_to_output_file ###########################################