47 lines
1.7 KiB
Python
47 lines
1.7 KiB
Python
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) |