bone_2026/src/data_loader.py

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)