levind_abhi

This commit is contained in:
Ubuntu
2025-05-23 10:17:25 +00:00
parent bec8344e93
commit a57e48c814
3 changed files with 83 additions and 31 deletions
+27 -15
View File
@@ -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)
+24 -6
View File
@@ -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)
+32 -10
View File
@@ -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