diff --git a/__init__.py b/__init__.py index 45a6f38..5746525 100644 --- a/__init__.py +++ b/__init__.py @@ -238,7 +238,7 @@ class TRI3DLEVINDABHICLOTHSEGBATCH: }, } - RETURN_TYPES = ("IMAGE", ) + RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE") FUNCTION = "main" CATEGORY = "TRI3D" @@ -292,21 +292,33 @@ class TRI3DLEVINDABHICLOTHSEGBATCH: # Collect and return the results - - - batch_results = [] + mask0_batch = [] + mask1_batch = [] + mask2_batch = [] for i in range(images.shape[0]): - cv2_segm = cv2.imread(LSEG_OUTPUT_PATH + f'image{i}.png', cv2.IMREAD_UNCHANGED) # Read PNG with alpha channel - cv2_segm = cv2.cvtColor(cv2_segm, cv2.COLOR_BGRA2RGBA) # Convert from BGRA to RGBA - b_tensor_img = cv2_img_to_tensor(cv2_segm) - batch_results.append(b_tensor_img.squeeze(0)) - - - batch_results = torch.stack(batch_results) - - - - return (batch_results, ) + mask0_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask0.png") + mask1_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask1.png") + mask2_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask2.png") + mask0_img = cv2.imread(mask0_path, cv2.IMREAD_UNCHANGED) + mask1_img = cv2.imread(mask1_path, cv2.IMREAD_UNCHANGED) + mask2_img = cv2.imread(mask2_path, cv2.IMREAD_UNCHANGED) + # Ensure single channel, convert to 3 channel if needed for consistency + if mask0_img is not None and len(mask0_img.shape) == 2: + mask0_img = cv2.cvtColor(mask0_img, cv2.COLOR_GRAY2RGB) + if mask1_img is not None and len(mask1_img.shape) == 2: + mask1_img = cv2.cvtColor(mask1_img, cv2.COLOR_GRAY2RGB) + if mask2_img is not None and len(mask2_img.shape) == 2: + mask2_img = cv2.cvtColor(mask2_img, cv2.COLOR_GRAY2RGB) + mask0_tensor = cv2_img_to_tensor(mask0_img).squeeze(0) + mask1_tensor = cv2_img_to_tensor(mask1_img).squeeze(0) + mask2_tensor = cv2_img_to_tensor(mask2_img).squeeze(0) + mask0_batch.append(mask0_tensor) + mask1_batch.append(mask1_tensor) + mask2_batch.append(mask2_tensor) + mask0_batch = torch.stack(mask0_batch) + mask1_batch = torch.stack(mask1_batch) + mask2_batch = torch.stack(mask2_batch) + return (mask0_batch, mask1_batch, mask2_batch) diff --git a/cloth-segmentation/app.py b/cloth-segmentation/app.py index 818b72a..5b0b04b 100644 --- a/cloth-segmentation/app.py +++ b/cloth-segmentation/app.py @@ -14,16 +14,34 @@ def initialize_and_load_models(): net = initialize_and_load_models() -def run(img): +def run(img, image_id, output_dir): palette = get_palette(4) - cloth_seg = generate_mask(img, net=net,device=device) - return cloth_seg + mask0, mask1, mask2, cloth_seg = generate_mask(img, net=net, device=device, image_id=image_id, output_dir=output_dir) + return mask0, mask1, mask2, cloth_seg INPUT_PATH = "./input/" OUTPUT_PATH = "./output/" import os -for cur_image in os.listdir(INPUT_PATH): +mask0_paths = [] +mask1_paths = [] +mask2_paths = [] +cloth_paths = [] + +for idx, cur_image in enumerate(os.listdir(INPUT_PATH)): img = PIL.Image.open(INPUT_PATH + cur_image) - cloth_seg = run(img) - cloth_seg.save(OUTPUT_PATH + cur_image, format="PNG") + mask0, mask1, mask2, cloth_seg = run(img, image_id=idx, output_dir=OUTPUT_PATH) + # Save masks and cloth_seg with unique names (already saved in generate_mask) + mask0_path = os.path.join(OUTPUT_PATH, f"{idx}__mask0.png") + mask1_path = os.path.join(OUTPUT_PATH, f"{idx}__mask1.png") + mask2_path = os.path.join(OUTPUT_PATH, f"{idx}__mask2.png") + cloth_path = os.path.join(OUTPUT_PATH, f"{idx}__extracted_garment.png") + mask0_paths.append(mask0_path) + mask1_paths.append(mask1_path) + mask2_paths.append(mask2_path) + cloth_paths.append(cloth_path) + +print("Mask0 batch:", mask0_paths) +print("Mask1 batch:", mask1_paths) +print("Mask2 batch:", mask2_paths) +print("Garment batch:", cloth_paths) diff --git a/cloth-segmentation/process.py b/cloth-segmentation/process.py index 4206dfb..0653622 100644 --- a/cloth-segmentation/process.py +++ b/cloth-segmentation/process.py @@ -101,14 +101,16 @@ def apply_transform(img): from PIL import Image -def generate_mask(input_image, net, device='cpu'): +def generate_mask(input_image, net, device='cpu', image_id=None, output_dir=None): img = input_image img_size = img.size img = img.resize((768, 768), Image.BICUBIC) image_tensor = apply_transform(img) image_tensor = torch.unsqueeze(image_tensor, 0) - output_dir = os.path.join(opt.output, 'extracted_garment') + # Allow output_dir override for batch processing + if output_dir is None: + output_dir = os.path.join(opt.output, 'extracted_garment') os.makedirs(output_dir, exist_ok=True) with torch.no_grad(): @@ -118,30 +120,50 @@ def generate_mask(input_image, net, device='cpu'): output_tensor = torch.squeeze(output_tensor, dim=0) output_arr = output_tensor.cpu().numpy() + # Create and save individual masks for classes 1, 2, 3 + classes_of_interest = [1, 2, 3] + mask_imgs = [] + for idx, cls in enumerate(classes_of_interest): + mask = np.zeros_like(output_arr, dtype=np.uint8) + mask[output_arr == cls] = 255 + if mask.ndim > 2: + mask = mask.squeeze() + if mask.ndim != 2: + raise ValueError(f"mask{idx} must be a 2-dimensional array") + mask_img = Image.fromarray(mask, mode='L').resize(img_size, Image.BICUBIC) + # Save with unique name if image_id is provided + if image_id is not None: + mask_path = os.path.join(output_dir, f'{image_id}__mask{idx}.png') + else: + mask_path = os.path.join(output_dir, f'mask{idx}.png') + mask_img.save(mask_path, format="PNG") + print(f"Saved mask{idx} at: {mask_path}") + mask_imgs.append(mask_img) + # Create a binary mask where selected classes are 1, others are 0 binary_mask = np.zeros_like(output_arr, dtype=np.uint8) - classes_of_interest = [1, 2, 3] # Modify this list according to your classes of interest for cls in classes_of_interest: binary_mask[output_arr == cls] = 255 - - # Ensure binary_mask is 2D if binary_mask.ndim > 2: - binary_mask = binary_mask.squeeze() # Removes single-dimensional entries from the shape + binary_mask = binary_mask.squeeze() if binary_mask.ndim != 2: raise ValueError("binary_mask must be a 2-dimensional array") - binary_mask_img = Image.fromarray(binary_mask, mode='L').resize(img_size, Image.BICUBIC) # Create an RGBA image for the output extracted_garment = Image.new("RGBA", img_size) - original_img = img.resize(img_size) # Resize the processed image back to original size + original_img = img.resize(img_size) extracted_garment.paste(original_img, mask=binary_mask_img) # Save the garment image with transparency - garment_path = os.path.join(output_dir, 'extracted_garment.png') + if image_id is not None: + garment_path = os.path.join(output_dir, f'{image_id}__extracted_garment.png') + else: + garment_path = os.path.join(output_dir, 'extracted_garment.png') extracted_garment.save(garment_path, format="PNG") + print(f"Saved extracted garment at: {garment_path}") - return extracted_garment + return (*mask_imgs, extracted_garment) # def generate_mask(input_image, net, device='cpu'): # img = input_image