From 060eeb3457ab0384006ede2d4e7ed1a233b829ef Mon Sep 17 00:00:00 2001 From: Radionic Date: Tue, 14 Nov 2023 18:29:10 +0800 Subject: [PATCH] chore: uncomment backend SAM node --- sam/sam_multilayer.py | 141 +++++++++++++++++++++--------------------- 1 file changed, 72 insertions(+), 69 deletions(-) diff --git a/sam/sam_multilayer.py b/sam/sam_multilayer.py index 8f1e2b2..1549f4d 100644 --- a/sam/sam_multilayer.py +++ b/sam/sam_multilayer.py @@ -43,80 +43,83 @@ class SAMMultiLayer: FUNCTION = "load_image" def load_image(self, image, ckpt, embedding_id, image_prompts_json): - # global global_predictor - # model_type = re.findall(r'vit_[lbh]', ckpt)[0] - - # if global_predictor is None: - # ckpt = folder_paths.get_full_path("sams", ckpt) - # sam = sam_model_registry[model_type](checkpoint=ckpt) - # predictor = SamPredictor(sam) - # global_predictor = predictor - - # predictor = global_predictor - - # emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy" - # if not os.path.exists(emb_filename): - # image_np = (image[0].numpy() * 255).astype(np.uint8) - # predictor.set_image(image_np) - # emb = predictor.get_image_embedding().cpu().numpy() - # np.save(emb_filename, emb) - - # with open(f"{self.output_dir}/{embedding_id}_{model_type}.json", "w") as f: - # data = { - # "input_size": predictor.input_size, - # "original_size": predictor.original_size, - # } - # json.dump(data, f) - # else: - # emb = np.load(emb_filename) - - # with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f: - # data = json.load(f) - # predictor.input_size = data["input_size"] - # predictor.features = torch.from_numpy(emb) - # predictor.is_image_set = True - # predictor.original_size = data["original_size"] - - # image_prompts = json.loads(image_prompts_json) - - # result = [image_prompts] - - # if isinstance(image_prompts, list): - # pass - # elif all(isinstance(item, list) for item in image_prompts.values()): - # for item in image_prompts.values(): - # if (len(item) == 0): - # h, w, c = image[0].shape - # result.append(torch.zeros(1, h, w, c)) - # continue - # point_coords = np.array([[p['x'], p['y']] for p in item]) - # point_labels = np.array([p['label'] for p in item]) - - # masks, _, _ = predictor.predict( - # point_coords=point_coords, - # point_labels=point_labels, - # ) - # masks = torch.from_numpy(masks) - # masks = rearrange(masks[0], 'h w -> 1 h w') - # out_image = repeat(masks, '1 h w -> 1 h w c', c=3) * image - # result.append(out_image) - image_prompts = json.loads(image_prompts_json) - result = [image_prompts] order_file = f"{self.output_dir}/{embedding_id}/order.json" - if not os.path.exists(order_file): - raise FileNotFoundError("Segments order file not found. Please click the 'Edit prompt' in the graph first.") + if os.path.exists(order_file): + # Frontend uploads segments images to backend => backend reads all segments images and passes them to next nodes + with open(order_file) as f: + order = json.load(f) - with open(order_file) as f: - order = json.load(f) + result = [image_prompts] - for segment in order: - image = Image.open(f"{self.output_dir}/{embedding_id}/{segment}.png") - image = np.array(image).astype(np.float32) / 255.0 - image = torch.from_numpy(image)[None,] - result.append(image) - return result + for segment in order: + image = Image.open(f"{self.output_dir}/{embedding_id}/{segment}.png") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + result.append(image) + + return result + else: + # Frontend uploads clicks coordinates to backend => backend runs SAM and passes the segments to next nodes + global global_predictor + model_type = re.findall(r'vit_[lbh]', ckpt)[0] + + if global_predictor is None: + ckpt = folder_paths.get_full_path("sams", ckpt) + sam = sam_model_registry[model_type](checkpoint=ckpt) + predictor = SamPredictor(sam) + global_predictor = predictor + + predictor = global_predictor + + emb_filename = f"{self.output_dir}/{embedding_id}_{model_type}.npy" + if not os.path.exists(emb_filename): + image_np = (image[0].numpy() * 255).astype(np.uint8) + predictor.set_image(image_np) + emb = predictor.get_image_embedding().cpu().numpy() + np.save(emb_filename, emb) + + with open(f"{self.output_dir}/{embedding_id}_{model_type}.json", "w") as f: + data = { + "input_size": predictor.input_size, + "original_size": predictor.original_size, + } + json.dump(data, f) + else: + emb = np.load(emb_filename) + + with open(f"{self.output_dir}/{embedding_id}_{model_type}.json") as f: + data = json.load(f) + predictor.input_size = data["input_size"] + predictor.features = torch.from_numpy(emb) + predictor.is_image_set = True + predictor.original_size = data["original_size"] + + image_prompts = json.loads(image_prompts_json) + + result = [image_prompts] + + if isinstance(image_prompts, list): + pass + elif all(isinstance(item, list) for item in image_prompts.values()): + for item in image_prompts.values(): + if (len(item) == 0): + h, w, c = image[0].shape + result.append(torch.zeros(1, h, w, c)) + continue + point_coords = np.array([[p['x'], p['y']] for p in item]) + point_labels = np.array([p['label'] for p in item]) + + masks, _, _ = predictor.predict( + point_coords=point_coords, + point_labels=point_labels, + ) + masks = torch.from_numpy(masks) + masks = rearrange(masks[0], 'h w -> 1 h w') + out_image = repeat(masks, '1 h w -> 1 h w c', c=3) * image + result.append(out_image) + return result NODE_CLASS_MAPPINGS = {"SAM MultiLayer": SAMMultiLayer}