mwmathis's picture C-Achard's picture
Add PyTorch SuperAnimal backend and refresh the Space (#14)
2e1b62d
Raw History Blame
10.5 kB
import json
import math
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"
###########################################
# JSON outputs: pixel coordinates in the input image, hidden keypoints as null
COORDINATES = "pixels in the input image (after EXIF orientation), origin top-left, x right, y down"
def keypoints_to_json(kpts, offset=(0, 0), scale=(1.0, 1.0)):
"""[x, y, confidence] per keypoint, mapped by (k + offset) * scale; NaN (below threshold) becomes null."""
out = []
for x, y, conf in np.asarray(kpts, dtype=float)[:, :3]:
if np.isnan(conf):
out.append(None)
else:
out.append([(x + offset[0]) * scale[0], (y + offset[1]) * scale[1], conf])
return out
def write_json(info, path_to_output_file):
with open(path_to_output_file, "w") as f:
json.dump(info, f, indent=1, allow_nan=False)
return path_to_output_file
def json_header(image_size, annotated_size):
return {
"date": str(today),
"coordinates": COORDINATES,
"image_size": list(image_size),
"annotated_image_size": list(annotated_size),
}
def save_results_as_json(
md_results,
dlc_outputs,
animal_bboxes,
map_dlc_label_id_to_str,
model,
mega_model_input,
image_size,
path_to_output_file="download_predictions.json",
):
"""TF (legacy) MegaDetector + DLC results.
animal_bboxes: detection rows [x1,y1,x2,y2,conf,label] in the MegaDetector frame, one per entry of dlc_outputs
dlc_outputs: keypoints relative to each crop (crops start at floor(x1), floor(y1), see crop_animal_detections)
"""
md_h, md_w = md_results.ims[0].shape[:2]
scale = (image_size[0] / md_w, image_size[1] / md_h)
info = json_header(image_size, (md_w, md_h))
info["MD_model"] = str(mega_model_input)
info["number_of_bb"] = len(dlc_outputs)
info["dlc_model"] = model
labels = list(map_dlc_label_id_to_str.values())
for i, kpts in enumerate(dlc_outputs):
x1, y1, x2, y2, confidence, _ = animal_bboxes[i]
info["bb_" + str(i)] = {
"corner_1": (x1 * scale[0], y1 * scale[1]),
"corner_2": (x2 * scale[0], y2 * scale[1]),
"predict MD": md_results.names[0],
"confidence MD": float(confidence),
"dlc_pred": dict(
zip(labels, keypoints_to_json(kpts, offset=(math.floor(x1), math.floor(y1)), scale=scale), strict=True)
),
}
return write_json(info, path_to_output_file)
def save_results_only_dlc(
dlc_outputs, map_label_id_to_str, model, image_size, output_file="dowload_predictions_dlc.json"
):
"""TF (legacy) DLC run on the whole input image (keypoints already in input pixels)."""
info = json_header(image_size, image_size)
info["dlc_model"] = model
info["dlc_pred"] = dict(zip(map_label_id_to_str.values(), keypoints_to_json(dlc_outputs), strict=True))
return write_json(info, output_file)
def save_results_pytorch(
animals,
map_label_id_to_str,
model,
pose_model,
detector,
image_size,
annotated_size,
path_to_output_file="download_predictions.json",
):
"""PyTorch SuperAnimal results (same layout as save_results_as_json).
animals: list of {'bbox': [x1,y1,x2,y2,conf], 'kpts': (num_keypoints, 3)}, in the annotated (resized) image
detector: None if the detector was skipped (whole image used as one animal)
"""
scale = (image_size[0] / annotated_size[0], image_size[1] / annotated_size[1])
info = json_header(image_size, annotated_size)
info["backend"] = "pytorch"
info["dlc_model"] = model
info["pose_model"] = pose_model
info["detector"] = detector
info["number_of_bb"] = len(animals)
labels = list(map_label_id_to_str.values())
for i, animal in enumerate(animals):
x1, y1, x2, y2, confidence = animal["bbox"]
info["bb_" + str(i)] = {
"corner_1": (x1 * scale[0], y1 * scale[1]),
"corner_2": (x2 * scale[0], y2 * scale[1]),
"confidence": confidence,
"dlc_pred": dict(zip(labels, keypoints_to_json(animal["kpts"], scale=scale), strict=True)),
}
return write_json(info, path_to_output_file)
###########################################