BioMike/clipsegmulticlass
0
1import os
2from PIL import Image
3import torch
4from torch.utils.data import Dataset
5import numpy as np
6
7class SingleClassSegmentationDataset(Dataset):
8 def __init__(self, dataset, class_labels, image_size=352, transform=None):
9
10 self.items = dataset
11 self.class_labels = class_labels
12 self.image_size = image_size
13 self.transform = transform
14
15 def __len__(self):
16 return len(self.items)
17
18 def __getitem__(self, idx):
19 item = self.items[idx]
20
21 image = Image.open(item["img_path"]).convert("RGB")
22 mask = Image.open(item["mask_path"]).convert("L")
23 class_name = item["label"]
24
25 class_index = self.class_labels.index(class_name)
26 background_index = 0
27
28 mask_np = np.array(mask) > 0
29 final_mask = np.full(mask_np.shape, background_index, dtype=np.uint8)
30 final_mask[mask_np] = class_index
31
32 image = image.resize((self.image_size, self.image_size), Image.BILINEAR)
33 final_mask = Image.fromarray(final_mask).resize((self.image_size, self.image_size), Image.NEAREST)
34
35 if self.transform:
36 image, final_mask = self.transform(image, final_mask)
37
38 return {
39 "image": image,
40 "labels": torch.from_numpy(np.array(final_mask)).long()
41 }
42
43
44class SegmentationCollator:
45 def __init__(self, processor, class_labels):
46 self.processor = processor
47 self.class_labels = class_labels
48
49 def __call__(self, batch):
50 images = [item["image"] for item in batch]
51 labels = [item["labels"] for item in batch]
52
53 prompts = self.class_labels * len(images)
54 expanded_images = [img for img in images for _ in self.class_labels]
55
56 inputs = self.processor(
57 images=expanded_images,
58 text=prompts,
59 return_tensors="pt",
60 padding=True,
61 truncation=True
62 )
63
64 return {
65 "pixel_values": inputs["pixel_values"],
66 "input_ids": inputs["input_ids"],
67 "labels": torch.stack(labels)
68 }
69 