From 877dd93b0ce87e3a893068f9c4fed32064a98936 Mon Sep 17 00:00:00 2001 From: Craig Wright Date: Wed, 26 Mar 2025 21:41:19 +0000 Subject: [PATCH 1/2] fix: Resolve tensor dimension mismatch in SegformerNode - Improve mask handling and preview generation - Add better error handling and dimension checks - Fix image and mask tensor format compatibility --- nodes/SegformerNode.py | 202 +++++++++++++++++++++++------------------ 1 file changed, 116 insertions(+), 86 deletions(-) diff --git a/nodes/SegformerNode.py b/nodes/SegformerNode.py index df277a4..b11d846 100644 --- a/nodes/SegformerNode.py +++ b/nodes/SegformerNode.py @@ -149,7 +149,29 @@ class SegformerNode: return mask def create_preview(self, image, mask): - # Create an RGBA preview with the mask as alpha channel + # Ensure image is in CHW format + if len(image.shape) == 2: + image = image.unsqueeze(0).repeat(3, 1, 1) + elif len(image.shape) == 3: + if image.shape[0] != 3: # If channels are not in first dimension + image = image.permute(2, 0, 1) # Move channels to first dimension + + # Ensure mask has correct dimensions + if len(mask.shape) == 3: + mask = mask.squeeze(0) + if len(mask.shape) > 2: + mask = mask.squeeze() + + # Resize mask to match image dimensions if needed + if mask.shape != image.shape[1:]: + mask = torch.nn.functional.interpolate( + mask.unsqueeze(0).unsqueeze(0), + size=image.shape[1:], + mode='bilinear', + align_corners=False + ).squeeze() + + # Create preview by concatenating image and mask preview = image.clone() preview = torch.cat([preview, mask.unsqueeze(0)], dim=0) return preview @@ -170,97 +192,105 @@ class SegformerNode: resize_mode="bilinear", invert_mask=False, show_preview=True, return_individual_masks=False, post_process="none", post_process_radius=3, segment_groups=""): - # Handle local checkpoint loading - if model_name.startswith("local:"): - local_dir = Path("models/segformer") / model_name[6:] - self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir) - else: - self.model, self.processor = SegformerModelLoader.load_model(model_name) + try: + # Handle local checkpoint loading + if model_name.startswith("local:"): + local_dir = Path("models/segformer") / model_name[6:] + self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir) + else: + self.model, self.processor = SegformerModelLoader.load_model(model_name) - show_on_node = False - - # Process input image - i = 255. * image[0].cpu().numpy() - img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) - inputs = self.processor(images=img, return_tensors="pt") - - # Get model outputs - outputs = self.model(**inputs) - logits = outputs.logits.cpu() - - # Upsample logits with specified resize mode - upsampled_logits = nn.functional.interpolate( - logits, - size=img.size[::-1], - mode=resize_mode, - align_corners=False if resize_mode != "nearest" else None, - ) - - pred_seg = upsampled_logits.argmax(dim=1)[0] - - # Parse segment groups if provided - segment_groups_dict = self.parse_segment_groups(segment_groups) - - # Create individual masks if requested - individual_masks = {} - segment_info = [] - - # Get unique segments and process each - unique_segments = np.unique(pred_seg.numpy()) - for segment in unique_segments: - segment_name = self.model.config.id2label[segment] - segment_info.append(f"Segment {segment}: {segment_name}") + show_on_node = False + # Process input image + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + inputs = self.processor(images=img, return_tensors="pt") + + # Get model outputs + outputs = self.model(**inputs) + logits = outputs.logits.cpu() + + # Upsample logits with specified resize mode + upsampled_logits = nn.functional.interpolate( + logits, + size=img.size[::-1], + mode=resize_mode, + align_corners=False if resize_mode != "nearest" else None, + ) + + pred_seg = upsampled_logits.argmax(dim=1)[0] + + # Parse segment groups if provided + segment_groups_dict = self.parse_segment_groups(segment_groups) + + # Create individual masks if requested + individual_masks = {} + segment_info = [] + + # Get unique segments and process each + unique_segments = np.unique(pred_seg.numpy()) + for segment in unique_segments: + segment_name = self.model.config.id2label[segment] + segment_info.append(f"Segment {segment}: {segment_name}") + + if return_individual_masks: + mask = (pred_seg == segment).float() + mask = self.process_mask(mask, normalize_mask, binary_mask, + invert_mask, post_process, post_process_radius) + individual_masks[segment_name] = mask + + # Create merged mask based on segment groups + if segment_groups_dict: + merged_mask = torch.zeros_like(pred_seg, dtype=torch.float32) + for group_name, indices in segment_groups_dict.items(): + group_mask = torch.zeros_like(pred_seg, dtype=torch.float32) + for idx in indices: + group_mask = torch.maximum(group_mask, (pred_seg == idx).float()) + merged_mask = torch.maximum(merged_mask, group_mask) + segment_info.append(f"Group {group_name}: {indices}") + else: + merged_mask = torch.ones_like(pred_seg, dtype=torch.float32) + + # Process the final mask + merged_mask = self.process_mask(merged_mask, normalize_mask, binary_mask, + invert_mask, post_process, post_process_radius) + + # Create visualization + fig = plt.figure() + plt.imshow(pred_seg) + buf = io.BytesIO() + plt.savefig(buf, format='png') + buf.seek(0) + img2 = Image.open(buf) + + # Convert visualization to tensor + i = ImageOps.exif_transpose(img2) + if i.getbands() != ("R", "G", "B", "A"): + i = i.convert("RGBA") + img2 = np.array(img2).astype(np.float32) / 255.0 + img2 = torch.from_numpy(img2)[None,] + + # Create preview if requested + preview = self.create_preview(image[0], merged_mask) if show_preview else None + + # Join segment info + segment_info_str = "\n".join(segment_info) if return_individual_masks: - mask = (pred_seg == segment).float() - mask = self.process_mask(mask, normalize_mask, binary_mask, - invert_mask, post_process, post_process_radius) - individual_masks[segment_name] = mask + segment_info_str += "\n\nIndividual masks available for: " + ", ".join(individual_masks.keys()) - # Create merged mask based on segment groups - if segment_groups_dict: - merged_mask = torch.zeros_like(pred_seg, dtype=torch.float32) - for group_name, indices in segment_groups_dict.items(): - group_mask = torch.zeros_like(pred_seg, dtype=torch.float32) - for idx in indices: - group_mask = torch.maximum(group_mask, (pred_seg == idx).float()) - merged_mask = torch.maximum(merged_mask, group_mask) - segment_info.append(f"Group {group_name}: {indices}") - else: - merged_mask = torch.ones_like(pred_seg, dtype=torch.float32) + output_ui = {"images": [img2]} if show_on_node else {} - # Process the final mask - merged_mask = self.process_mask(merged_mask, normalize_mask, binary_mask, - invert_mask, post_process, post_process_radius) + # Return results + return {"result": (img2, merged_mask, segment_info_str, preview if preview is not None else img2), + "ui": output_ui} - # Create visualization - fig = plt.figure() - plt.imshow(pred_seg) - buf = io.BytesIO() - plt.savefig(buf, format='png') - buf.seek(0) - img2 = Image.open(buf) - - # Convert visualization to tensor - i = ImageOps.exif_transpose(img2) - if i.getbands() != ("R", "G", "B", "A"): - i = i.convert("RGBA") - img2 = np.array(img2).astype(np.float32) / 255.0 - img2 = torch.from_numpy(img2)[None,] - - # Create preview if requested - preview = self.create_preview(image[0], merged_mask) if show_preview else None - - # Join segment info - segment_info_str = "\n".join(segment_info) - if return_individual_masks: - segment_info_str += "\n\nIndividual masks available for: " + ", ".join(individual_masks.keys()) - - output_ui = {"images": [img2]} if show_on_node else {} - - # Return results - return {"result": (img2, merged_mask, segment_info_str, preview if preview is not None else img2), - "ui": output_ui} + except Exception as e: + import traceback + print(f"Error in segmentation: {str(e)}") + print(f"Traceback: {traceback.format_exc()}") + return {"result": (image, torch.zeros_like(image[0, :, :, 0]), f"Error: {str(e)}", image), + "ui": {"images": [image]} if show_on_node else {}} class SegformerNodeMasks: @classmethod From 8fb9b27ea1d130caae96a0a7060be2f84623d190 Mon Sep 17 00:00:00 2001 From: Craig Wright Date: Wed, 26 Mar 2025 21:44:41 +0000 Subject: [PATCH 2/2] Update nodes/SegformerNode.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- nodes/SegformerNode.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/nodes/SegformerNode.py b/nodes/SegformerNode.py index b11d846..2950e97 100644 --- a/nodes/SegformerNode.py +++ b/nodes/SegformerNode.py @@ -289,7 +289,7 @@ class SegformerNode: import traceback print(f"Error in segmentation: {str(e)}") print(f"Traceback: {traceback.format_exc()}") - return {"result": (image, torch.zeros_like(image[0, :, :, 0]), f"Error: {str(e)}", image), + return {"result": (image, torch.zeros_like(image[0, :, :]), f"Error: {str(e)}", image), "ui": {"images": [image]} if show_on_node else {}} class SegformerNodeMasks: