Набор данных с двумя входами
Построение модели с несколькими входами начинается с создания пользовательского набора данных, который будет передавать все входные данные модели. В этом упражнении вы реализуете набор данных Omniglot, формирующий тройки из следующих элементов:
- Изображение символа, который нужно классифицировать,
- Вектор алфавита в формате унитарного кодирования (one-hot encoding) длиной 30, где все значения равны нулю, кроме одного, указывающего идентификатор алфавита, к которому принадлежит символ,
- Целевая метка — целое число от 0 до 963.
Вам предоставлен список samples, состоящий из 3-элементных кортежей: путь к файлу изображения, вектор алфавита и целевая метка. Следующие импорты уже выполнены, так что можно приступать!
from PIL import Image
from torch.utils.data import DataLoader, Dataset
from torchvision import transforms
Это упражнение является частью курса
Глубокое обучение на PyTorch: средний уровень
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
class OmniglotDataset(Dataset):
def __init__(self, transform, samples):
# Assign transform and samples to class attributes
____ = ____
____ = ____