Метод __len__ класса Dataset
Метод __len__ класса Dataset возвращает общее
количество элементов (образцов) в пользовательском наборе данных.
Этот метод необходим для корректной работы загрузчиков данных
DataLoader, которые используют его для определения длины
датасета и организации итераций по данным. Метод не принимает
никаких параметров и должен возвращать целое неотрицательное число.
Синтаксис
class CustomDataset(Dataset):
def __len__(self):
return len(self.data)
Пример
Создадим простой датасет с фиксированным набором данных и
реализуем метод __len__:
import torch
from torch.utils.data import Dataset
class SimpleDataset(Dataset):
def __init__(self):
self.data = [1, 2, 3, 4, 5]
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
dataset = SimpleDataset()
print(len(dataset))
Результат выполнения кода:
5
Пример
Рассмотрим датасет с динамически генерируемыми данными, где размер известен заранее:
import torch
from torch.utils.data import Dataset
class GeneratedDataset(Dataset):
def __init__(self, length, feature_size):
self.length = length
self.feature_size = feature_size
def __len__(self):
return self.length
def __getitem__(self, idx):
return torch.randn(self.feature_size)
torch.manual_seed(0)
dataset = GeneratedDataset(10, 5)
print(len(dataset))
Результат выполнения кода:
10
Пример
Покажем, как метод __len__ используется в цикле
для перебора датасета:
import torch
from torch.utils.data import Dataset
class WordDataset(Dataset):
def __init__(self, words):
self.words = words
def __len__(self):
return len(self.words)
def __getitem__(self, idx):
return self.words[idx]
dataset = WordDataset(["apple", "banana", "cherry"])
print(f"Size: {len(dataset)}")
for i in range(len(dataset)):
print(f"Item {i}: {dataset[i]}")
Результат выполнения кода:
Size: 3
Item 0: apple
Item 1: banana
Item 2: cherry
Смотрите также
-
класс
Dataset,
который является базовым для всех пользовательских датасетов -
метод
__getitem__,
который возвращает элемент датасета по индексу -
метод
__add__,
который позволяет объединять два датасета в один -
метод
__len__,
который возвращает количество элементов в датасете