levind_abhi
This commit is contained in:
+27
-15
@@ -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)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user