Team Ai
Apppublic

BioMike/clipsegmulticlass

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
data_processing.py69 linesDownload Raw Back to src
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