From 96ec5db7dad429cb56cf843431058090668d85aa Mon Sep 17 00:00:00 2001 From: Jordan Phillips Date: Sat, 13 Jul 2024 07:37:32 -0700 Subject: [PATCH] Make output sort order configurable --- nodes.py | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index 01028e3..2185dfa 100644 --- a/nodes.py +++ b/nodes.py @@ -26,14 +26,15 @@ def prepare_deepface_home(): os.environ["DEEPFACE_HOME"] = deepface_path -def result_from_images_with_distances(image_tuples): - image_tuples.sort(key=lambda row: row[1]) - images = [row[0] for row in image_tuples] - distances = [row[1] for row in image_tuples] - verified_ratios = [row[2] for row in image_tuples] +def result_from_images_with_measurements(images_with_measurements, sort_by): + images_with_measurements.sort(key=lambda row: row[1] if sort_by == "distance" else -row[2]) + + images = [row[0] for row in images_with_measurements] + distances = [row[1] for row in images_with_measurements] + ratios = [row[2] for row in images_with_measurements] if len(images) > 0: - return torch.stack(images, dim=0), distances, verified_ratios + return torch.stack(images, dim=0), distances, ratios else: return (None, None, None,) @@ -136,6 +137,12 @@ class DeepfaceVerifyNode: ], { "default": "VGG-Face", }), + "sort_by": ([ + "distance", + "ratio" + ], { + "default": "distance", + }) }, } @@ -153,7 +160,7 @@ class DeepfaceVerifyNode: CATEGORY = "deepface" - def run(self, images, reference_images, distance_threshold, ratio_threshold, detector_backend, model_name): + def run(self, images, reference_images, distance_threshold, ratio_threshold, detector_backend, model_name, sort_by): deepface_reference_images = [] for reference_image in reference_images: deepface_reference_images.append(deepface_image_from_comfy_image(reference_image)) @@ -207,7 +214,7 @@ class DeepfaceVerifyNode: image_counter += 1 - return result_from_images_with_distances(verified_image_tuples) + result_from_images_with_distances(rejected_image_tuples) + return result_from_images_with_measurements(verified_image_tuples, sort_by) + result_from_images_with_measurements(rejected_image_tuples, sort_by) NODE_CLASS_MAPPINGS = { "DeepfaceExtractFaces": DeepfaceExtractFacesNode,