Функция default_convert
Функция default_convert является частью модуля torch.data
и используется для преобразования различных типов данных в тензоры PyTorch.
Она автоматически определяет тип входных данных и конвертирует их в соответствующий
тензор. Функция часто применяется внутри пользовательских загрузчиков данных
для приведения данных к единому тензорному формату.
Функция default_convert принимает один обязательный параметр -
входные данные. Она может обрабатывать скаляры, списки, кортежи, массивы NumPy
и другие структуры данных. Результатом работы функции всегда является тензор
PyTorch или структура из тензоров.
Синтаксис
torch.utils.data.default_convert(data)
Параметры функции:
-
data- входные данные для преобразования. Может быть скаляром, списком, кортежем, массивом NumPy или их вложенными структурами.
Функция возвращает тензор PyTorch или структуру, состоящую из тензоров (например, список тензоров, кортеж тензоров или словарь тензоров).
Пример
Давайте преобразуем список чисел в тензор с помощью default_convert:
import torch
from torch.utils.data import default_convert
data = [1, 2, 3, 4, 5]
t = default_convert(data)
print(t)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5])
Пример
Функция default_convert автоматически обрабатывает вложенные структуры данных,
такие как кортеж списков или список кортежей:
import torch
from torch.utils.data import default_convert
data = ([1, 2, 3], [4, 5, 6])
t = default_convert(data)
print(t)
Результат выполнения кода:
(tensor([1, 2, 3]), tensor([4, 5, 6]))
Пример
Давайте преобразуем словарь с данными разных типов с помощью default_convert.
Функция сохраняет структуру словаря, преобразуя значения в тензоры:
import torch
from torch.utils.data import default_convert
data = {
'features': [1, 2, 3],
'labels': [0, 1, 0],
'scalar': 42
}
res = default_convert(data)
print(res['features'])
print(res['labels'])
print(res['scalar'])
Результат выполнения кода:
tensor([1, 2, 3])
tensor([0, 1, 0])
tensor(42)
Пример
Функция default_convert автоматически определяет тип данных элементов.
Например, если в списке есть числа с плавающей точкой, результатом будет тензор
с типом float:
import torch
from torch.utils.data import default_convert
data = [1, 2.5, 3, 4.7]
t = default_convert(data)
print(t)
print(t.dtype)
Результат выполнения кода:
tensor([1.0000, 2.5000, 3.0000, 4.7000])
torch.float32
Пример
Давайте используем default_convert внутри пользовательского класса
для загрузки данных. Это позволяет автоматически преобразовывать данные
в тензоры при обращении к элементам датасета:
import torch
from torch.utils.data import Dataset, default_convert
class MyDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
return default_convert(item)
raw_data = [
([1, 2], 0),
([3, 4], 1),
([5, 6], 0)
]
dataset = MyDataset(raw_data)
first_item = dataset[0]
print(first_item)
Результат выполнения кода:
(tensor([1, 2]), tensor(0))
Смотрите также
-
функцию
default_collate,
которая объединяет список выборок в батч тензоров -
класс
TensorDataset,
который оборачивает тензоры в датасет для удобной загрузки -
функцию
get_worker_info,
которая возвращает информацию о процессе-загрузчике данных -
класс
Subset,
который создает подмножество данных из датасета по индексам