import torch from torch.utils.data import Dataset import numpy as np from PIL import Image import os class FlexibleDataset(Dataset): """Датасет, который можно адаптировать под любые данные""" def __init__(self, data_dir, is_train=True, input_size=(512, 512)): self.data_dir = data_dir self.input_size = input_size self.is_train = is_train # Собираем все изображения self.images = [] self.masks = [] # Ищем файлы (эта логика будет меняться под DICOM) for f in os.listdir(os.path.join(data_dir, 'images')): if f.endswith(('.jpg', '.png')): self.images.append(os.path.join(data_dir, 'images', f)) mask_path = os.path.join(data_dir, 'masks', f.replace('.jpg', '.png')) if os.path.exists(mask_path): self.masks.append(mask_path) def __len__(self): return len(self.images) def __getitem__(self, idx): # Загрузка изображения image = Image.open(self.images[idx]).convert('RGB') image = image.resize(self.input_size) image = np.array(image).transpose(2, 0, 1) / 255.0 # Загрузка маски (если есть) if idx < len(self.masks): mask = Image.open(self.masks[idx]) mask = mask.resize(self.input_size) mask = np.array(mask) # Бинаризация (0 - фон, 1 - объект) mask = (mask > 128).astype(np.int64) else: mask = np.zeros(self.input_size, dtype=np.int64) return torch.FloatTensor(image), torch.LongTensor(mask)