Функция get_worker_info
Функция get_worker_info возвращает объект с информацией о текущем процессе-воркере, который используется для параллельной загрузки данных в DataLoader.
Если функция вызвана вне процесса-воркера (в главном процессе), она возвращает None.
Объект информации содержит идентификатор воркера, общее количество воркеров и другие полезные данные для распределения нагрузки между воркерами.
Функция принимает только один необязательный параметр worker_id, но обычно вызывается без аргументов.
Синтаксис
torch.utils.data.get_worker_info()
Пример
Давайте проверим, вызывается ли функция в процессе-воркере:
import torch
from torch.utils.data import DataLoader, Dataset
class SimpleDataset(Dataset):
def __len__(self):
return 10
def __getitem__(self, idx):
info = torch.utils.data.get_worker_info()
if info is not None:
worker_id = info.id
return torch.tensor([idx, worker_id])
return torch.tensor([idx, -1])
dataset = SimpleDataset()
dataloader = DataLoader(dataset, batch_size=2, num_workers=2)
for batch in dataloader:
print(batch)
break
Результат выполнения кода:
tensor([
[0, 0],
[1, 0],
])
В примере видно, что оба элемента из первого батча обработаны воркером с идентификатором 0.
Пример
Теперь используем информацию о воркере для распределения данных. Каждый воркер будет обрабатывать свою часть датасета:
import torch
from torch.utils.data import DataLoader, Dataset
class PartitionDataset(Dataset):
def __init__(self, data):
self.data = data
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
return self.data[idx]
data = list(range(100))
dataset = PartitionDataset(data)
def worker_init_fn(worker_id):
info = torch.utils.data.get_worker_info()
if info is not None:
worker_data = data[worker_id::info.num_workers]
info.dataset.data = worker_data
dataloader = DataLoader(
dataset,
batch_size=4,
num_workers=2,
worker_init_fn=worker_init_fn
)
for batch in dataloader:
print(batch)
break
Результат выполнения кода:
tensor([0, 2, 4, 6])
В этом примере каждый воркер получает свою часть данных. Воркер 0 обрабатывает четные индексы, воркер 1 - нечетные.
Пример
Проверим возвращаемое значение функции при вызове вне воркера:
import torch
info = torch.utils.data.get_worker_info()
print(info)
Результат выполнения кода:
None
Смотрите также
-
класс
DataLoader,
который используется для загрузки данных с поддержкой параллельной обработки -
класс
Dataset,
который представляет набор данных для загрузки -
функцию
random_split,
которая разбивает датасет на подмножества -
класс
Subset,
который создает подмножество датасета по заданным индексам